mod adapt;
mod adaptive;
mod error;
mod images;
mod metrics;
mod protocol;
mod rate_limit;
mod request_log;
mod routes;
mod stream;
mod throttle;
use std::{
net::IpAddr,
num::{NonZeroU64, NonZeroUsize},
time::Duration,
};
use clap::{CommandFactory, FromArgMatches, Parser};
use dbx_tools_core::{build_info, init_logging_with_verbose, shutdown_signal, DatabricksClient};
use dbx_tools_model::{ModelCapabilitiesResolver, ModelClient, ModelRateLimitsResolver};
use images::DEFAULT_IMAGE_RESIZE_THRESHOLD_BYTES;
use metrics::{default_metrics_option, MetricsConfig, MetricsOption, MetricsRuntime, PeerAddr};
use protocol::TargetWire;
use rate_limit::{ModelFallbackMode, ModelFallbackPolicy, RateLimitPolicy};
use routes::{AppConfig, AppState};
use throttle::{RateLimitMode, ThrottleConfig};
use tracing::info;
const DEFAULT_MAX_REQUEST_BYTES: NonZeroUsize =
NonZeroUsize::new(25_000_000).expect("default request limit is non-zero");
const DEFAULT_IMAGE_RESIZE_THRESHOLD: NonZeroUsize =
NonZeroUsize::new(DEFAULT_IMAGE_RESIZE_THRESHOLD_BYTES)
.expect("default image resize threshold is non-zero");
const DEFAULT_RATE_LIMIT_RETRIES: u32 = 5;
const DEFAULT_RATE_LIMIT_INITIAL_DELAY_MS: NonZeroU64 =
NonZeroU64::new(1_000).expect("default retry delay is non-zero");
const DEFAULT_RATE_LIMIT_MAX_DELAY_MS: NonZeroU64 =
NonZeroU64::new(60_000).expect("default maximum retry delay is non-zero");
const DEFAULT_RATE_LIMIT_MAX_WAIT_MS: u64 = 60_000;
const DEFAULT_RATE_LIMIT_MODEL_FALLBACK_MAX_STEPS: u32 = 5;
const DEFAULT_RATE_LIMIT_MODEL_FALLBACK_THRESHOLD_MS: NonZeroU64 =
NonZeroU64::new(10_000).expect("default model fallback threshold is non-zero");
#[derive(Debug, Parser)]
#[command(name = "dbx-model-proxy")]
struct Cli {
#[arg(short = 'v', long)]
verbose: bool,
#[arg(long, env = "DATABRICKS_CONFIG_PROFILE")]
profile: Option<String>,
#[arg(long, default_value = "127.0.0.1")]
host: IpAddr,
#[arg(long, env = "DATABRICKS_APP_PORT", default_value_t = 4000)]
port: u16,
#[arg(long, env = "METRICS", default_value_t = default_metrics_option())]
metrics: MetricsOption,
#[arg(long, env = "METRICS_PUBLIC", default_value_t = false)]
metrics_public: bool,
#[arg(long, value_enum, default_value_t = TargetWire::Auto)]
target: TargetWire,
#[arg(long, env = "MAX_REQUEST_BYTES", default_value_t = DEFAULT_MAX_REQUEST_BYTES)]
max_request_bytes: NonZeroUsize,
#[arg(
long,
env = "IMAGE_RESIZE_THRESHOLD_BYTES",
default_value_t = DEFAULT_IMAGE_RESIZE_THRESHOLD
)]
image_resize_threshold_bytes: NonZeroUsize,
#[arg(long, env = "INPUT_TOKENS_PER_MINUTE")]
input_tokens_per_minute: Option<NonZeroU64>,
#[arg(long, env = "OUTPUT_TOKENS_PER_MINUTE")]
output_tokens_per_minute: Option<NonZeroU64>,
#[arg(long, env = "PROVISIONED_THROUGHPUT", default_value_t = false)]
provisioned_throughput: bool,
#[arg(long, env = "RATE_LIMIT_MODE", value_enum, default_value_t = RateLimitMode::Auto)]
rate_limit_mode: RateLimitMode,
#[arg(
long,
env = "RATE_LIMIT_MODEL_FALLBACK",
value_enum,
default_value_t = ModelFallbackMode::SameFamily
)]
rate_limit_model_fallback: ModelFallbackMode,
#[arg(
long,
env = "RATE_LIMIT_MODEL_FALLBACK_MAX_STEPS",
default_value_t = DEFAULT_RATE_LIMIT_MODEL_FALLBACK_MAX_STEPS,
value_parser = clap::value_parser!(u32).range(..=20)
)]
rate_limit_model_fallback_max_steps: u32,
#[arg(
long,
env = "RATE_LIMIT_MODEL_FALLBACK_THRESHOLD_MS",
default_value_t = DEFAULT_RATE_LIMIT_MODEL_FALLBACK_THRESHOLD_MS
)]
rate_limit_model_fallback_threshold_ms: NonZeroU64,
#[arg(
long,
env = "RATE_LIMIT_MAX_WAIT_MS",
default_value_t = DEFAULT_RATE_LIMIT_MAX_WAIT_MS,
value_parser = clap::value_parser!(u64).range(1..=60_000)
)]
rate_limit_max_wait_ms: u64,
#[arg(
long,
env = "RATE_LIMIT_RETRIES",
default_value_t = DEFAULT_RATE_LIMIT_RETRIES,
value_parser = clap::value_parser!(u32).range(..=100)
)]
rate_limit_retries: u32,
#[arg(
long,
env = "RATE_LIMIT_INITIAL_DELAY_MS",
default_value_t = DEFAULT_RATE_LIMIT_INITIAL_DELAY_MS
)]
rate_limit_initial_delay_ms: NonZeroU64,
#[arg(
long,
env = "RATE_LIMIT_MAX_DELAY_MS",
default_value_t = DEFAULT_RATE_LIMIT_MAX_DELAY_MS
)]
rate_limit_max_delay_ms: NonZeroU64,
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let cli = Cli::from_arg_matches(&Cli::command().version(build_info::version()).get_matches())?;
init_logging_with_verbose(cli.verbose)?;
let Cli {
verbose: _,
profile,
host,
port,
metrics,
metrics_public,
target,
max_request_bytes,
image_resize_threshold_bytes,
input_tokens_per_minute,
output_tokens_per_minute,
provisioned_throughput,
rate_limit_mode,
rate_limit_model_fallback,
rate_limit_model_fallback_max_steps,
rate_limit_model_fallback_threshold_ms,
rate_limit_max_wait_ms,
rate_limit_retries,
rate_limit_initial_delay_ms,
rate_limit_max_delay_ms,
} = cli;
if rate_limit_initial_delay_ms > rate_limit_max_delay_ms {
return Err("RATE_LIMIT_INITIAL_DELAY_MS must not exceed RATE_LIMIT_MAX_DELAY_MS".into());
}
let metrics = MetricsRuntime::new(MetricsConfig::resolve(metrics, host, metrics_public)?)?;
let databricks = DatabricksClient::new(profile).await?;
let models = ModelClient::new(databricks.clone())?;
let model_rate_limits = if provisioned_throughput || rate_limit_mode == RateLimitMode::Off {
Default::default()
} else {
ModelRateLimitsResolver::new()?.rate_limits().await?
};
let state = AppState::new(
ModelCapabilitiesResolver::new()?,
databricks,
models,
AppConfig {
target,
throttle: ThrottleConfig {
input_tokens_per_minute,
output_tokens_per_minute,
provisioned_throughput,
mode: rate_limit_mode,
documented_limits: model_rate_limits,
},
image_resize_threshold_bytes: image_resize_threshold_bytes.get(),
model_fallback: ModelFallbackPolicy {
mode: rate_limit_model_fallback,
max_steps: rate_limit_model_fallback_max_steps as usize,
threshold: Duration::from_millis(rate_limit_model_fallback_threshold_ms.get()),
},
rate_limits: RateLimitPolicy {
max_retries: rate_limit_retries,
initial_delay: Duration::from_millis(rate_limit_initial_delay_ms.get()),
max_delay: Duration::from_millis(rate_limit_max_delay_ms.get()),
max_wait: Duration::from_millis(rate_limit_max_wait_ms),
},
metrics: metrics.clone(),
},
);
let listener = tokio::net::TcpListener::bind((host, port)).await?;
let address = listener.local_addr()?;
info!(
address = %address,
?target,
metrics_mode = %metrics.mode(),
metrics_routes_visible = metrics.routes_visible(),
metrics_public,
max_request_bytes = max_request_bytes.get(),
image_resize_threshold_bytes = image_resize_threshold_bytes.get(),
input_tokens_per_minute = input_tokens_per_minute.map(NonZeroU64::get),
output_tokens_per_minute = output_tokens_per_minute.map(NonZeroU64::get),
provisioned_throughput,
?rate_limit_mode,
?rate_limit_model_fallback,
rate_limit_model_fallback_max_steps,
rate_limit_model_fallback_threshold_ms = rate_limit_model_fallback_threshold_ms.get(),
rate_limit_max_wait_ms,
rate_limit_retries,
rate_limit_initial_delay_ms = rate_limit_initial_delay_ms.get(),
rate_limit_max_delay_ms = rate_limit_max_delay_ms.get(),
"model proxy listening"
);
let listener = metrics.track_listener(listener);
axum::serve(
listener,
routes::routes(state, max_request_bytes).into_make_service_with_connect_info::<PeerAddr>(),
)
.with_graceful_shutdown(shutdown_signal())
.await?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn request_limit_defaults_to_image_capable_proxy_limit() {
let cli = Cli::try_parse_from(["dbx-model-proxy"]).unwrap();
assert_eq!(cli.max_request_bytes, DEFAULT_MAX_REQUEST_BYTES);
assert!(!cli.verbose);
assert_eq!(cli.metrics, default_metrics_option());
assert!(!cli.metrics_public);
assert_eq!(cli.max_request_bytes.get(), 25_000_000);
assert_eq!(
cli.image_resize_threshold_bytes,
DEFAULT_IMAGE_RESIZE_THRESHOLD
);
assert_eq!(cli.input_tokens_per_minute, None);
assert_eq!(cli.output_tokens_per_minute, None);
assert!(!cli.provisioned_throughput);
assert_eq!(cli.rate_limit_mode, RateLimitMode::Auto);
assert_eq!(cli.rate_limit_model_fallback, ModelFallbackMode::SameFamily);
assert_eq!(
cli.rate_limit_model_fallback_max_steps,
DEFAULT_RATE_LIMIT_MODEL_FALLBACK_MAX_STEPS
);
assert_eq!(
cli.rate_limit_model_fallback_threshold_ms,
DEFAULT_RATE_LIMIT_MODEL_FALLBACK_THRESHOLD_MS
);
assert_eq!(cli.rate_limit_max_wait_ms, DEFAULT_RATE_LIMIT_MAX_WAIT_MS);
assert_eq!(cli.rate_limit_retries, DEFAULT_RATE_LIMIT_RETRIES);
assert_eq!(cli.rate_limit_retries, 5);
assert_eq!(
cli.rate_limit_initial_delay_ms,
DEFAULT_RATE_LIMIT_INITIAL_DELAY_MS
);
assert_eq!(cli.rate_limit_max_delay_ms, DEFAULT_RATE_LIMIT_MAX_DELAY_MS);
let cli = Cli::try_parse_from([
"dbx-model-proxy",
"--max-request-bytes",
"8388608",
"--metrics",
"false",
"--metrics-public",
"--image-resize-threshold-bytes",
"3145728",
"--input-tokens-per-minute",
"200000",
"--output-tokens-per-minute",
"20000",
"--provisioned-throughput",
"--rate-limit-mode",
"off",
"--rate-limit-model-fallback",
"off",
"--rate-limit-model-fallback-max-steps",
"3",
"--rate-limit-model-fallback-threshold-ms",
"7500",
"--rate-limit-max-wait-ms",
"30000",
"--rate-limit-retries",
"0",
"--rate-limit-initial-delay-ms",
"250",
"--rate-limit-max-delay-ms",
"5000",
])
.unwrap();
assert_eq!(cli.max_request_bytes.get(), 8 * 1024 * 1024);
assert_eq!(cli.metrics, MetricsOption::Off);
assert!(cli.metrics_public);
assert_eq!(cli.image_resize_threshold_bytes.get(), 3 * 1024 * 1024);
assert_eq!(cli.input_tokens_per_minute.unwrap().get(), 200_000);
assert_eq!(cli.output_tokens_per_minute.unwrap().get(), 20_000);
assert!(cli.provisioned_throughput);
assert_eq!(cli.rate_limit_mode, RateLimitMode::Off);
assert_eq!(cli.rate_limit_model_fallback, ModelFallbackMode::Off);
assert_eq!(cli.rate_limit_model_fallback_max_steps, 3);
assert_eq!(cli.rate_limit_model_fallback_threshold_ms.get(), 7_500);
assert_eq!(cli.rate_limit_max_wait_ms, 30_000);
assert_eq!(cli.rate_limit_retries, 0);
assert_eq!(cli.rate_limit_initial_delay_ms.get(), 250);
assert_eq!(cli.rate_limit_max_delay_ms.get(), 5_000);
assert!(
Cli::try_parse_from(["dbx-model-proxy", "--rate-limit-max-wait-ms", "60001",]).is_err()
);
}
}