diff --git a/.github/workflows/docker.yml b/.github/workflows/docker.yml index b83c483..cefb8da 100644 --- a/.github/workflows/docker.yml +++ b/.github/workflows/docker.yml @@ -11,7 +11,7 @@ on: - main env: - IMAGE_NAME: ${{ secrets.DOCKERHUB_USERNAME }}/anyllm-proxy + IMAGE_NAME: followthewhit3rabbit/anyllm-proxy jobs: test: @@ -72,6 +72,16 @@ jobs: - uses: docker/setup-buildx-action@v4 + - name: Validate Docker Hub credentials + env: + DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_USERNAME }} + DOCKERHUB_TOKEN: ${{ secrets.DOCKERHUB_TOKEN }} + run: | + if [ -z "${DOCKERHUB_USERNAME}" ] || [ -z "${DOCKERHUB_TOKEN}" ]; then + echo "DOCKERHUB_USERNAME and DOCKERHUB_TOKEN secrets are required to publish Docker images." + exit 1 + fi + - name: Log in to Docker Hub uses: docker/login-action@v4 with: @@ -87,7 +97,7 @@ jobs: type=semver,pattern={{version}} type=semver,pattern={{major}}.{{minor}} type=sha,prefix=sha-,format=short - type=raw,value=latest,enable=true + type=raw,value=latest,enable=${{ startsWith(github.ref, 'refs/tags/v') && !contains(github.ref_name, '-') }} - name: Build and push by digest id: build @@ -100,6 +110,8 @@ jobs: cache-from: type=gha,scope=${{ matrix.platform }} cache-to: type=gha,scope=${{ matrix.platform }},mode=max labels: ${{ steps.meta.outputs.labels }} + provenance: mode=max + sbom: true - name: Export digest run: | @@ -135,6 +147,16 @@ jobs: - uses: docker/setup-buildx-action@v4 + - name: Validate Docker Hub credentials + env: + DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_USERNAME }} + DOCKERHUB_TOKEN: ${{ secrets.DOCKERHUB_TOKEN }} + run: | + if [ -z "${DOCKERHUB_USERNAME}" ] || [ -z "${DOCKERHUB_TOKEN}" ]; then + echo "DOCKERHUB_USERNAME and DOCKERHUB_TOKEN secrets are required to publish Docker images." + exit 1 + fi + - name: Log in to Docker Hub uses: docker/login-action@v4 with: @@ -150,7 +172,7 @@ jobs: type=semver,pattern={{version}} type=semver,pattern={{major}}.{{minor}} type=sha,prefix=sha-,format=short - type=raw,value=latest,enable=true + type=raw,value=latest,enable=${{ startsWith(github.ref, 'refs/tags/v') && !contains(github.ref_name, '-') }} - name: Create and push multi-arch manifest working-directory: /tmp/digests @@ -168,3 +190,10 @@ jobs: docker buildx imagetools create \ "${tag_args[@]}" \ "${digest_args[@]}" + + - name: Verify multi-arch manifest + run: | + image_tag=$(jq -r '.tags[0]' <<< '${{ steps.meta.outputs.json }}') + docker buildx imagetools inspect "${image_tag}" > manifest.txt + grep -q 'linux/amd64' manifest.txt + grep -q 'linux/arm64' manifest.txt diff --git a/README.md b/README.md index ce49a2a..97ad942 100644 --- a/README.md +++ b/README.md @@ -453,6 +453,14 @@ cp .env.example .env # set OPENAI_API_KEY docker compose up ``` +Published images are on [Docker Hub](https://hub.docker.com/r/followthewhit3rabbit/anyllm-proxy). CI publishes multi-arch images for `linux/amd64` and `linux/arm64` when a `v*` release tag is pushed. + +Release tags: +- `X.Y.Z` for the exact release +- `X.Y` for the latest patch in a minor series +- `sha-` for the release commit +- `latest` for stable `vX.Y.Z` releases only, prerelease tags do not update it +
Smoke tests (no real API key needed) diff --git a/crates/client/src/lib.rs b/crates/client/src/lib.rs index a579bc7..e7ca635 100644 --- a/crates/client/src/lib.rs +++ b/crates/client/src/lib.rs @@ -85,7 +85,7 @@ pub use retry::{ backoff_delay, is_quota_exhausted, is_retryable, parse_retry_after, send_with_retry, send_with_retry_policy, RetryPolicy, RetryableError, }; -pub use sse::{find_double_newline, SseError}; +pub use sse::{find_double_newline, SseError, SseFrameBuffer}; pub use tools::{ToolBuilder, ToolChoiceBuilder}; // Re-export key types from the translator crate so downstream users diff --git a/crates/client/src/sse.rs b/crates/client/src/sse.rs index 3a9ee09..6ffdc9a 100644 --- a/crates/client/src/sse.rs +++ b/crates/client/src/sse.rs @@ -4,7 +4,7 @@ //! boundaries (`\n\n` or `\r\n\r\n`), and delivers each `data:` line to a //! caller-supplied callback. No dependency on axum or any web framework. -use bytes::BytesMut; +use bytes::{Bytes, BytesMut}; /// Maximum SSE buffer size (10 MB). Protects against unbounded memory growth /// if the backend sends data without frame delimiters. @@ -42,6 +42,60 @@ pub fn find_double_newline(buf: &[u8], start: usize) -> Option<(usize, usize)> { None } +/// Bounded SSE frame buffer. +/// +/// It owns the partial byte buffer for a stream, enforces a maximum buffered +/// size before appending new bytes, and returns complete frame byte slices +/// without copying frame contents. +pub struct SseFrameBuffer { + buffer: BytesMut, + search_from: usize, + max_size: usize, +} + +impl SseFrameBuffer { + /// Create a frame buffer using the default SSE maximum. + pub fn new() -> Self { + Self::with_max_size(MAX_SSE_BUFFER_SIZE) + } + + /// Create a frame buffer with an explicit maximum, useful for tests. + pub fn with_max_size(max_size: usize) -> Self { + Self { + buffer: BytesMut::new(), + search_from: 0, + max_size, + } + } + + /// Append bytes and return all complete frames found in the buffer. + pub fn push(&mut self, bytes: &[u8]) -> Result, SseError> { + let Some(new_len) = self.buffer.len().checked_add(bytes.len()) else { + return Err(SseError::BufferOverflow); + }; + if new_len > self.max_size { + return Err(SseError::BufferOverflow); + } + self.buffer.extend_from_slice(bytes); + + let mut frames = Vec::new(); + while let Some((pos, delim_len)) = find_double_newline(&self.buffer, self.search_from) { + let frame = self.buffer.split_to(pos).freeze(); + let _ = self.buffer.split_to(delim_len); + self.search_from = 0; + frames.push(frame); + } + self.search_from = self.buffer.len().saturating_sub(3); + Ok(frames) + } +} + +impl Default for SseFrameBuffer { + fn default() -> Self { + Self::new() + } +} + /// Read SSE frames from a response stream, calling `on_data` for each `data:` line. /// /// Returns `Ok(())` on normal stream completion, or an `SseError` on failure. @@ -63,21 +117,14 @@ where use futures::StreamExt; let mut stream = response.bytes_stream(); // BytesMut (not String) because TCP chunks may split mid-UTF-8 character. - let mut buffer = BytesMut::new(); + let mut buffer = SseFrameBuffer::new(); let mut frame_events: Vec = Vec::new(); - let mut search_from: usize = 0; while let Some(chunk_result) = stream.next().await { let bytes = chunk_result?; - buffer.extend_from_slice(&bytes); - - if buffer.len() > MAX_SSE_BUFFER_SIZE { - return Err(SseError::BufferOverflow); - } - - while let Some((pos, delim_len)) = find_double_newline(&buffer, search_from) { + for frame in buffer.push(&bytes)? { frame_events.clear(); - match std::str::from_utf8(&buffer[..pos]) { + match std::str::from_utf8(&frame) { Ok(frame_str) => { for line in frame_str.lines() { let line = line.trim(); @@ -92,16 +139,11 @@ where tracing::warn!("skipping non-UTF-8 SSE frame: {e}"); } } - let _ = buffer.split_to(pos + delim_len); - search_from = 0; if !on_events(&frame_events) { return Ok(()); // consumer disconnected } } - // Next chunk: resume scanning 3 bytes back from the end. The 4-byte - // delimiter \r\n\r\n could straddle the chunk boundary. - search_from = buffer.len().saturating_sub(3); } Ok(()) @@ -157,4 +199,31 @@ mod tests { assert_eq!(pos, 0); assert_eq!(len, 2); } + + #[test] + fn frame_buffer_extracts_complete_frames_across_chunks() { + let mut buffer = SseFrameBuffer::new(); + + assert!(buffer.push(b"data: one\r\n").unwrap().is_empty()); + let frames = buffer + .push(b"\r\ndata: two\n\ndata: partial") + .expect("valid frame chunks"); + + assert_eq!(frames.len(), 2); + assert_eq!(&frames[0][..], b"data: one"); + assert_eq!(&frames[1][..], b"data: two"); + } + + #[test] + fn frame_buffer_rejects_oversized_chunk_before_append() { + let mut buffer = SseFrameBuffer::with_max_size(4); + + assert!(buffer.push(b"12").unwrap().is_empty()); + let err = buffer.push(b"345").unwrap_err(); + + assert!(matches!(err, SseError::BufferOverflow)); + let frames = buffer.push(b"\n\n").unwrap(); + assert_eq!(frames.len(), 1); + assert_eq!(&frames[0][..], b"12"); + } } diff --git a/crates/client/src/streaming.rs b/crates/client/src/streaming.rs index e227a8a..1bba36b 100644 --- a/crates/client/src/streaming.rs +++ b/crates/client/src/streaming.rs @@ -4,12 +4,11 @@ use anyllm_translate::anthropic::streaming::StreamEvent; use anyllm_translate::mapping; use anyllm_translate::openai::ChatCompletionChunk; -use bytes::BytesMut; use futures::{SinkExt, Stream, StreamExt}; use pin_project_lite::pin_project; use crate::error::ClientError; -use crate::sse::{find_double_newline, SseError, MAX_SSE_BUFFER_SIZE}; +use crate::sse::{SseError, SseFrameBuffer}; /// Argument to the [`run_sse_task`] handler. Either a parsed UTF-8 SSE frame /// or the end-of-stream signal (bytes exhausted without a transport error). @@ -30,8 +29,7 @@ async fn run_sse_task( mut handler: impl FnMut(SseEvent<'_>) -> Vec>, ) { let mut stream = response.bytes_stream(); - let mut buffer = BytesMut::new(); - let mut search_from = 0usize; + let mut buffer = SseFrameBuffer::new(); while let Some(chunk_result) = stream.next().await { let bytes = match chunk_result { @@ -41,25 +39,22 @@ async fn run_sse_task( return; } }; - buffer.extend_from_slice(&bytes); + let frames = match buffer.push(&bytes) { + Ok(frames) => frames, + Err(e) => { + let _ = tx.send(Err(ClientError::Sse(e))).await; + return; + } + }; - if buffer.len() > MAX_SSE_BUFFER_SIZE { - let _ = tx - .send(Err(ClientError::Sse(SseError::BufferOverflow))) - .await; - return; - } - - while let Some((pos, delim_len)) = find_double_newline(&buffer, search_from) { - let events = match std::str::from_utf8(&buffer[..pos]) { + for frame in frames { + let events = match std::str::from_utf8(&frame) { Ok(frame_str) => handler(SseEvent::Frame(frame_str)), Err(e) => { tracing::warn!("skipping non-UTF-8 SSE frame: {e}"); vec![] } }; - let _ = buffer.split_to(pos + delim_len); - search_from = 0; for event in events { if tx.send(event).await.is_err() { @@ -67,7 +62,6 @@ async fn run_sse_task( } } } - search_from = buffer.len().saturating_sub(3); } for event in handler(SseEvent::End) { diff --git a/crates/providers/src/catalog.rs b/crates/providers/src/catalog.rs index 0e804d9..c0c7ba9 100644 --- a/crates/providers/src/catalog.rs +++ b/crates/providers/src/catalog.rs @@ -58,6 +58,8 @@ pub struct ProviderCatalog { providers: BTreeMap, advertised_provider_ids: BTreeSet, models_by_provider: BTreeMap>, + provider_ids_by_litellm_prefix: BTreeMap, + model_indexes_by_provider: BTreeMap>, } #[derive(Debug)] @@ -220,6 +222,8 @@ impl ProviderCatalog { providers, advertised_provider_ids, models_by_provider, + provider_ids_by_litellm_prefix: BTreeMap::new(), + model_indexes_by_provider: BTreeMap::new(), }; catalog.refresh_metadata_counts(); catalog @@ -325,9 +329,12 @@ impl ProviderCatalog { } pub fn get_model(&self, provider_id: &str, model_id: &str) -> Option<&OwnedModelDef> { - self.list_models(provider_id) - .iter() - .find(|model| model.id == model_id) + let provider_id = registry::canonical_provider_id(provider_id); + let index = self + .model_indexes_by_provider + .get(provider_id)? + .get(model_id)?; + self.models_by_provider.get(provider_id)?.get(*index) } pub fn resolve_backend(&self, provider_id: &str) -> Option<(&'static str, &str)> { @@ -346,12 +353,8 @@ impl ProviderCatalog { } pub fn find_by_litellm_prefix(&self, prefix: &str) -> Option<&OwnedProviderDef> { - let direct = self - .providers - .values() - .find(|p| !p.litellm_prefix.is_empty() && prefix == p.litellm_prefix); - if direct.is_some() { - return direct; + if let Some(provider_id) = self.provider_ids_by_litellm_prefix.get(prefix) { + return self.providers.get(provider_id); } let provider_id = prefix.strip_suffix('/')?; @@ -360,6 +363,7 @@ impl ProviderCatalog { } fn refresh_metadata_counts(&mut self) { + self.rebuild_indexes(); self.metadata.provider_count = self.all_providers().count(); self.metadata.model_count = self .advertised_provider_ids @@ -368,6 +372,26 @@ impl ProviderCatalog { .map(Vec::len) .sum(); } + + fn rebuild_indexes(&mut self) { + self.provider_ids_by_litellm_prefix.clear(); + for (provider_id, provider) in &self.providers { + if !provider.litellm_prefix.is_empty() { + self.provider_ids_by_litellm_prefix + .insert(provider.litellm_prefix.clone(), provider_id.clone()); + } + } + + self.model_indexes_by_provider.clear(); + for (provider_id, models) in &self.models_by_provider { + let mut indexes = BTreeMap::new(); + for (index, model) in models.iter().enumerate() { + indexes.insert(model.id.clone(), index); + } + self.model_indexes_by_provider + .insert(provider_id.clone(), indexes); + } + } } #[derive(Default)] diff --git a/crates/providers/src/catalog/tests.rs b/crates/providers/src/catalog/tests.rs index c66ba12..8723953 100644 --- a/crates/providers/src/catalog/tests.rs +++ b/crates/providers/src/catalog/tests.rs @@ -1,6 +1,6 @@ use super::*; use crate::model::ModelStatus; -use crate::provider::{ProviderCapabilities, ProviderStatus}; +use crate::provider::ProviderStatus; const TEST_FIXTURE: &str = r#"{ "openai/gpt-fresh": { diff --git a/crates/providers/src/registry.rs b/crates/providers/src/registry.rs index 797ad25..f2a0b2e 100644 --- a/crates/providers/src/registry.rs +++ b/crates/providers/src/registry.rs @@ -1,6 +1,8 @@ use crate::model::ModelDef; use crate::provider::{ProviderDef, ProviderProtocol}; use crate::providers; +use std::collections::HashMap; +use std::sync::LazyLock; /// All registered LiteLLM-compatible providers. /// @@ -93,6 +95,60 @@ static LEGACY_ONLY_MODELS: &[(&str, &[ModelDef])] = &[ ("xinference", providers::xinference::MODELS), ]; +static PROVIDERS_BY_ID: LazyLock> = + LazyLock::new(|| { + let mut map = HashMap::with_capacity(ALL_PROVIDERS.len() + LEGACY_ONLY_PROVIDERS.len()); + for provider in ALL_PROVIDERS + .iter() + .copied() + .chain(LEGACY_ONLY_PROVIDERS.iter().copied()) + { + map.insert(provider.id, provider); + } + map + }); + +static PROVIDERS_BY_LITELLM_PREFIX: LazyLock> = + LazyLock::new(|| { + let mut map = HashMap::new(); + for provider in ALL_PROVIDERS + .iter() + .copied() + .chain(LEGACY_ONLY_PROVIDERS.iter().copied()) + { + if !provider.litellm_prefix.is_empty() { + map.insert(provider.litellm_prefix, provider); + } + } + map + }); + +static MODELS_BY_PROVIDER: LazyLock> = + LazyLock::new(|| { + let mut map = HashMap::with_capacity(ALL_MODELS.len() + LEGACY_ONLY_MODELS.len()); + for (provider_id, models) in ALL_MODELS.iter().chain(LEGACY_ONLY_MODELS.iter()) { + map.insert(*provider_id, *models); + } + map + }); + +static MODEL_BY_PROVIDER_AND_ID: LazyLock< + HashMap<(&'static str, &'static str), &'static ModelDef>, +> = LazyLock::new(|| { + let model_count = ALL_MODELS + .iter() + .chain(LEGACY_ONLY_MODELS.iter()) + .map(|(_, models)| models.len()) + .sum(); + let mut map = HashMap::with_capacity(model_count); + for (provider_id, models) in ALL_MODELS.iter().chain(LEGACY_ONLY_MODELS.iter()) { + for model in *models { + map.insert((*provider_id, model.id), model); + } + } + map +}); + #[cfg(feature = "runtime-catalog")] pub(crate) fn advertised_provider_defs() -> &'static [&'static ProviderDef] { ALL_PROVIDERS @@ -125,11 +181,7 @@ pub fn canonical_provider_id(id: &str) -> &str { /// Look up a provider by its `id` field (e.g. `"groq"`, `"together_ai"`). pub fn get_provider(id: &str) -> Option<&'static ProviderDef> { let id = canonical_provider_id(id); - ALL_PROVIDERS - .iter() - .find(|p| p.id == id) - .copied() - .or_else(|| LEGACY_ONLY_PROVIDERS.iter().find(|p| p.id == id).copied()) + PROVIDERS_BY_ID.get(id).copied() } /// All registered providers. @@ -140,22 +192,15 @@ pub fn all_providers() -> impl Iterator { /// All models registered for a given provider id. pub fn list_models(provider_id: &str) -> &'static [ModelDef] { let provider_id = canonical_provider_id(provider_id); - ALL_MODELS - .iter() - .find(|(id, _)| *id == provider_id) - .map(|(_, models)| *models) - .or_else(|| { - LEGACY_ONLY_MODELS - .iter() - .find(|(id, _)| *id == provider_id) - .map(|(_, models)| *models) - }) - .unwrap_or(&[]) + MODELS_BY_PROVIDER.get(provider_id).copied().unwrap_or(&[]) } /// Look up a specific model by provider id and model id. pub fn get_model(provider_id: &str, model_id: &str) -> Option<&'static ModelDef> { - list_models(provider_id).iter().find(|m| m.id == model_id) + let provider_id = canonical_provider_id(provider_id); + MODEL_BY_PROVIDER_AND_ID + .get(&(provider_id, model_id)) + .copied() } /// Whether a native Anthropic model supports LiteLLM's adaptive-thinking mode. @@ -210,23 +255,13 @@ pub fn resolve_backend(provider_id: &str) -> Option<(&'static str, &'static str) /// Find a provider by its LiteLLM routing prefix (e.g. `"groq/"` or `"together_ai/"`). /// Used by `parse_provider_model()` in litellm config parsing. pub fn find_by_litellm_prefix(prefix: &str) -> Option<&'static ProviderDef> { - let direct = ALL_PROVIDERS - .iter() - .find(|p| !p.litellm_prefix.is_empty() && prefix == p.litellm_prefix) - .copied(); - if direct.is_some() { - return direct; + if let Some(provider) = PROVIDERS_BY_LITELLM_PREFIX.get(prefix).copied() { + return Some(provider); } let provider_id = prefix.strip_suffix('/')?; let canonical = canonical_provider_id(provider_id); - if canonical == provider_id { - return LEGACY_ONLY_PROVIDERS - .iter() - .find(|p| !p.litellm_prefix.is_empty() && prefix == p.litellm_prefix) - .copied(); - } - ALL_PROVIDERS.iter().find(|p| p.id == canonical).copied() + PROVIDERS_BY_ID.get(canonical).copied() } #[cfg(test)] diff --git a/crates/proxy/src/backend/mod.rs b/crates/proxy/src/backend/mod.rs index 6fd9dea..22964db 100644 --- a/crates/proxy/src/backend/mod.rs +++ b/crates/proxy/src/backend/mod.rs @@ -18,7 +18,7 @@ pub use anyllm_client::rate_limit::RateLimitHeaders; pub use anyllm_client::retry::{ backoff_delay, is_retryable, parse_retry_after, RetryableError, MAX_RETRIES, }; -pub use anyllm_client::sse::{find_double_newline, MAX_SSE_BUFFER_SIZE}; +pub use anyllm_client::sse::{find_double_newline, SseFrameBuffer, MAX_SSE_BUFFER_SIZE}; use anyllm_client::http::HttpClientConfig; diff --git a/crates/proxy/src/cache/mod.rs b/crates/proxy/src/cache/mod.rs index 95f7314..f379e8e 100644 --- a/crates/proxy/src/cache/mod.rs +++ b/crates/proxy/src/cache/mod.rs @@ -17,7 +17,7 @@ pub mod semantic; use bytes::Bytes; use sha2::{Digest, Sha256}; -use std::collections::BTreeMap; +use std::io::{self, Write}; use std::time::Instant; /// Maximum allowed value for per-request `cache_ttl_secs`. @@ -80,53 +80,95 @@ pub fn cache_key_for_request( ns: CacheNamespace, scope: &CacheScope<'_>, ) -> String { - // Fields that affect the backend response. Order does not matter because - // BTreeMap sorts keys alphabetically before serialization. - const CACHE_FIELDS: &[&str] = &[ - "_header_anthropic-beta", - "_header_x-claude-code-session-id", - "cache_ttl_secs", - "max_completion_tokens", - "max_tokens", - "messages", - "model", - "reasoning_effort", - "response_format", - "system", - "stop", - "temperature", - "tool_choice", - "tools", - "top_p", - ]; - - let mut canonical = BTreeMap::new(); - if let Some(obj) = body.as_object() { - for &field in CACHE_FIELDS { - if let Some(val) = obj.get(field) { - // Skip null values so absent fields and explicit null produce the same key. - if !val.is_null() { - canonical.insert(field, val.clone()); - } - } - } - } - canonical.insert( - "_scope_backend", - serde_json::Value::String(scope.backend_name.to_string()), - ); - canonical.insert( - "_scope_auth", - serde_json::Value::String(scope.auth_identity.to_string()), - ); - - // serde_json serializes BTreeMap in key order, giving us canonical JSON. - let json = serde_json::to_string(&canonical).unwrap_or_default(); - let hash = Sha256::digest(json.as_bytes()); + let mut hasher = Sha256::new(); + write_canonical_cache_body(&mut hasher, body, scope); + let hash = hasher.finalize(); let hex = hex::encode(hash); format!("{}:{}", ns.prefix(), hex) } +enum CacheField<'a> { + Json(&'a str, &'a serde_json::Value), + Str(&'static str, &'a str), +} + +impl<'a> CacheField<'a> { + fn key(&self) -> &str { + match self { + Self::Json(key, _) | Self::Str(key, _) => key, + } + } +} + +struct HashWriter<'a> { + hasher: &'a mut Sha256, +} + +impl Write for HashWriter<'_> { + fn write(&mut self, buf: &[u8]) -> io::Result { + self.hasher.update(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +fn write_canonical_cache_body( + hasher: &mut Sha256, + body: &serde_json::Value, + scope: &CacheScope<'_>, +) { + let mut fields = Vec::new(); + if let Some(obj) = body.as_object() { + fields.extend( + obj.iter() + .filter(|(key, value)| should_include_cache_field(key, value)) + .map(|(key, value)| CacheField::Json(key.as_str(), value)), + ); + } + fields.push(CacheField::Str("_scope_auth", scope.auth_identity)); + fields.push(CacheField::Str("_scope_backend", scope.backend_name)); + fields.sort_unstable_by(|a, b| a.key().cmp(b.key())); + + let mut writer = HashWriter { hasher }; + writer + .write_all(b"{") + .expect("hash writer should not fail writing object start"); + for (idx, field) in fields.iter().enumerate() { + if idx > 0 { + writer + .write_all(b",") + .expect("hash writer should not fail writing separator"); + } + serde_json::to_writer(&mut writer, field.key()) + .expect("hash writer should not fail writing key"); + writer + .write_all(b":") + .expect("hash writer should not fail writing colon"); + match field { + CacheField::Json(_, value) => serde_json::to_writer(&mut writer, value) + .expect("hash writer should not fail writing JSON value"), + CacheField::Str(_, value) => serde_json::to_writer(&mut writer, value) + .expect("hash writer should not fail writing string value"), + } + } + writer + .write_all(b"}") + .expect("hash writer should not fail writing object end"); +} + +fn should_include_cache_field(key: &str, value: &serde_json::Value) -> bool { + if value.is_null() { + return false; + } + !matches!( + key, + "stream" | "stream_options" | "_scope_auth" | "_scope_backend" + ) +} + pub struct CacheScope<'a> { pub backend_name: &'a str, pub auth_identity: &'a str, @@ -478,4 +520,66 @@ mod tests { "different cache_ttl_secs must produce different cache keys" ); } + + fn test_scope() -> CacheScope<'static> { + CacheScope { + backend_name: "openai", + auth_identity: "k1", + } + } + + fn anthropic_key(body: &serde_json::Value) -> String { + cache_key_for_request(body, CacheNamespace::Anthropic, &test_scope()) + } + + fn openai_key(body: &serde_json::Value) -> String { + cache_key_for_request(body, CacheNamespace::OpenAI, &test_scope()) + } + + #[test] + fn cache_key_includes_anthropic_response_affecting_fields() { + let base = serde_json::json!({ + "model": "claude-sonnet-4-6", + "max_tokens": 128, + "messages": [{"role": "user", "content": "hi"}] + }); + + let with_top_k = serde_json::json!({ + "model": "claude-sonnet-4-6", + "max_tokens": 128, + "messages": [{"role": "user", "content": "hi"}], + "top_k": 10 + }); + let with_stop_sequences = serde_json::json!({ + "model": "claude-sonnet-4-6", + "max_tokens": 128, + "messages": [{"role": "user", "content": "hi"}], + "stop_sequences": ["END"] + }); + let with_thinking = serde_json::json!({ + "model": "claude-sonnet-4-6", + "max_tokens": 128, + "messages": [{"role": "user", "content": "hi"}], + "thinking": {"type": "enabled", "budget_tokens": 1024} + }); + + assert_ne!(anthropic_key(&base), anthropic_key(&with_top_k)); + assert_ne!(anthropic_key(&base), anthropic_key(&with_stop_sequences)); + assert_ne!(anthropic_key(&base), anthropic_key(&with_thinking)); + } + + #[test] + fn cache_key_includes_unknown_extra_fields() { + let base = serde_json::json!({ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}] + }); + let with_extra = serde_json::json!({ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "prediction": {"type": "content", "content": "expected"} + }); + + assert_ne!(openai_key(&base), openai_key(&with_extra)); + } } diff --git a/crates/proxy/src/config/litellm/mod.rs b/crates/proxy/src/config/litellm/mod.rs index 066ab62..3fac1c8 100644 --- a/crates/proxy/src/config/litellm/mod.rs +++ b/crates/proxy/src/config/litellm/mod.rs @@ -10,6 +10,7 @@ use indexmap::IndexMap; use serde::Deserialize; use super::model_router::{Deployment, ModelRouter, RoutingStrategy}; +use super::single::validate_gcp_identifier; use super::{ resolve_env_value, validate_base_url, BackendAuth, BackendConfig, BackendKind, ModelMapping, MultiConfig, OpenAIApiFormat, TlsConfig, @@ -46,6 +47,9 @@ struct LiteLLMParams { weight: Option, // Azure-specific api_version: Option, + // Vertex-specific + vertex_project: Option, + vertex_location: Option, // Bedrock-specific aws_access_key_id: Option, aws_secret_access_key: Option, @@ -282,7 +286,7 @@ pub fn parse_litellm_yaml(yaml: &str) -> LiteLLMParsed { }), ); - let base_url = resolve_base_url(&kind, params, stub_provider); + let base_url = resolve_base_url(&kind, params, stub_provider, &actual_model); let bk = BackendKey { kind: format!("{kind:?}"), @@ -297,7 +301,15 @@ pub fn parse_litellm_yaml(yaml: &str) -> LiteLLMParsed { backend_counter += 1; let bc = build_backend_config( - &name, &kind, &api_key, &base_url, params, &tls, log_bodies, &config, + &name, + &kind, + &api_key, + &base_url, + &actual_model, + params, + &tls, + log_bodies, + &config, ); backend_map.insert(bk, (name.clone(), bc)); name @@ -415,10 +427,19 @@ fn resolve_base_url( kind: &BackendKind, params: &LiteLLMParams, stub_provider: Option<&'static anyllm_providers::ProviderDef>, + actual_model: &str, ) -> String { if let Some(ref url) = params.api_base { let resolved = resolve_env_value(url).unwrap_or_else(|e| panic!("model_list api_base: {e}")); + if *kind == BackendKind::AzureOpenAI && !resolved.contains("/openai/deployments/") { + let api_version = params.api_version.as_deref().unwrap_or("2024-10-21"); + let deployment = azure_deployment_from_model(actual_model); + return format!( + "{}/openai/deployments/{deployment}/chat/completions?api-version={api_version}", + resolved.trim_end_matches('/'), + ); + } return resolved; } match kind { @@ -457,11 +478,32 @@ fn resolve_base_url( panic!("api_base is required for azure deployments in model_list") } BackendKind::Vertex => { - panic!("api_base is required for vertex deployments in model_list") + let project = params.vertex_project.as_deref().unwrap_or_else(|| { + panic!("vertex_project is required for vertex deployments in model_list") + }); + let location = params.vertex_location.as_deref().unwrap_or_else(|| { + panic!("vertex_location is required for vertex deployments in model_list") + }); + validate_gcp_identifier("vertex_project", project); + validate_gcp_identifier("vertex_location", location); + format!( + "https://{location}-aiplatform.googleapis.com/v1/projects/{project}/locations/{location}/endpoints/openapi" + ) } } } +fn azure_deployment_from_model(model: &str) -> &str { + for marker in ["o_series/", "gpt5_series/"] { + if let Some(deployment) = model.strip_prefix(marker) { + if !deployment.is_empty() { + return deployment; + } + } + } + model +} + /// Build a BackendConfig from LiteLLM model_list params. #[allow(clippy::too_many_arguments)] fn build_backend_config( @@ -469,6 +511,7 @@ fn build_backend_config( kind: &BackendKind, api_key: &str, base_url: &str, + actual_model: &str, params: &LiteLLMParams, tls: &TlsConfig, log_bodies: bool, @@ -490,9 +533,10 @@ fn build_backend_config( // Already a full deployment URL. base_url.to_string() } else { + let deployment = azure_deployment_from_model(actual_model); format!( - "{}/openai/deployments/chat/completions?api-version={api_version}", - base_url.trim_end_matches('/') + "{}/openai/deployments/{deployment}/chat/completions?api-version={api_version}", + base_url.trim_end_matches('/'), ) } } else { diff --git a/crates/proxy/src/config/litellm/tests.rs b/crates/proxy/src/config/litellm/tests.rs index b4341fb..3709bfd 100644 --- a/crates/proxy/src/config/litellm/tests.rs +++ b/crates/proxy/src/config/litellm/tests.rs @@ -109,6 +109,131 @@ model_list: ); } +#[test] +fn azure_api_base_uses_model_deployment_name() { + let yaml = r#" +model_list: + - model_name: gpt-35 + litellm_params: + model: azure/chatgpt-v-2 + api_key: sk-azure + api_base: https://openai-gpt-4-test-v-1.openai.azure.com/ + api_version: "2023-05-15" +"#; + + let (multi, router) = from_litellm_yaml(yaml); + let bc = multi.backends.values().next().unwrap(); + assert_eq!( + bc.base_url, + "https://openai-gpt-4-test-v-1.openai.azure.com/openai/deployments/chatgpt-v-2/chat/completions?api-version=2023-05-15" + ); + assert_eq!(router.route("gpt-35").unwrap().actual_model, "chatgpt-v-2"); +} + +#[test] +fn azure_route_marker_is_not_used_as_deployment_name() { + let yaml = r#" +model_list: + - model_name: o3-mini + litellm_params: + model: azure/o_series/my-o3-deployment + api_key: sk-azure + api_base: https://azure-o-series.openai.azure.com +"#; + + let (multi, router) = from_litellm_yaml(yaml); + let bc = multi.backends.values().next().unwrap(); + assert_eq!( + bc.base_url, + "https://azure-o-series.openai.azure.com/openai/deployments/my-o3-deployment/chat/completions?api-version=2024-10-21" + ); + assert_eq!( + router.route("o3-mini").unwrap().actual_model, + "o_series/my-o3-deployment" + ); +} + +#[test] +fn azure_full_deployment_url_is_preserved() { + let yaml = r#" +model_list: + - model_name: gpt-4o + litellm_params: + model: azure/gpt-4o-deploy + api_key: sk-azure + api_base: https://myresource.openai.azure.com/openai/deployments/gpt-4o-deploy/chat/completions?api-version=2024-10-21 +"#; + + let (multi, _) = from_litellm_yaml(yaml); + let bc = multi.backends.values().next().unwrap(); + assert_eq!( + bc.base_url, + "https://myresource.openai.azure.com/openai/deployments/gpt-4o-deploy/chat/completions?api-version=2024-10-21" + ); +} + +#[test] +fn azure_deployments_on_same_resource_are_distinct_backends() { + let yaml = r#" +model_list: + - model_name: gpt-4o + litellm_params: + model: azure/gpt-4o-deploy + api_key: sk-azure + api_base: https://myresource.openai.azure.com + - model_name: gpt-4o-mini + litellm_params: + model: azure/gpt-4o-mini-deploy + api_key: sk-azure + api_base: https://myresource.openai.azure.com +"#; + + let (multi, router) = from_litellm_yaml(yaml); + assert_eq!(multi.backends.len(), 2); + + let first = router.route("gpt-4o").unwrap(); + let second = router.route("gpt-4o-mini").unwrap(); + assert_ne!(first.backend_name, second.backend_name); + + assert!(multi + .backends + .get(first.backend_name) + .unwrap() + .base_url + .contains("/openai/deployments/gpt-4o-deploy/")); + assert!(multi + .backends + .get(second.backend_name) + .unwrap() + .base_url + .contains("/openai/deployments/gpt-4o-mini-deploy/")); +} + +#[test] +fn vertex_litellm_project_and_location_build_base_url() { + let yaml = r#" +model_list: + - model_name: gemini-pro + litellm_params: + model: vertex_ai/gemini-2.5-pro + api_key: ya29-test + vertex_project: project-123 + vertex_location: us-central1 +"#; + + let (multi, router) = from_litellm_yaml(yaml); + let bc = multi.backends.values().next().unwrap(); + assert_eq!(bc.kind, BackendKind::Vertex); + assert_eq!( + bc.base_url, + "https://us-central1-aiplatform.googleapis.com/v1/projects/project-123/locations/us-central1/endpoints/openapi" + ); + assert_eq!( + router.route("gemini-pro").unwrap().actual_model, + "gemini-2.5-pro" + ); +} + #[test] fn multiple_deployments_same_model() { let yaml = r#" diff --git a/crates/proxy/src/config/single.rs b/crates/proxy/src/config/single.rs index e0f8a3f..d14c457 100644 --- a/crates/proxy/src/config/single.rs +++ b/crates/proxy/src/config/single.rs @@ -53,7 +53,7 @@ impl Config { let backend = match backend_str.to_ascii_lowercase().as_str() { "openai" => BackendKind::OpenAI, "azure" => BackendKind::AzureOpenAI, - "vertex" => BackendKind::Vertex, + "vertex" | "vertex_ai" => BackendKind::Vertex, "gemini" => BackendKind::Gemini, "anthropic" => BackendKind::Anthropic, "bedrock" => BackendKind::Bedrock, @@ -424,4 +424,33 @@ mod tests { clear_anthropic_env(); } + + #[test] + fn backend_vertex_ai_alias_uses_vertex_config() { + let _lock = crate::config::ENV_TEST_LOCK + .lock() + .unwrap_or_else(|e| e.into_inner()); + clear_anthropic_env(); + unsafe { + std::env::set_var("BACKEND", "vertex_ai"); + std::env::set_var("VERTEX_PROJECT", "project-123"); + std::env::set_var("VERTEX_REGION", "us-central1"); + std::env::set_var("VERTEX_API_KEY", "AIzaSy-test"); + std::env::remove_var("GOOGLE_ACCESS_TOKEN"); + } + + let config = Config::from_env(); + assert_eq!(config.backend, BackendKind::Vertex); + assert_eq!( + config.openai_base_url, + "https://us-central1-aiplatform.googleapis.com/v1/projects/project-123/locations/us-central1/endpoints/openapi" + ); + + unsafe { + std::env::remove_var("BACKEND"); + std::env::remove_var("VERTEX_PROJECT"); + std::env::remove_var("VERTEX_REGION"); + std::env::remove_var("VERTEX_API_KEY"); + } + } } diff --git a/crates/proxy/src/cost/mod.rs b/crates/proxy/src/cost/mod.rs index 3ed8e2f..1f5a989 100644 --- a/crates/proxy/src/cost/mod.rs +++ b/crates/proxy/src/cost/mod.rs @@ -6,6 +6,7 @@ pub mod db; use dashmap::DashMap; +use std::collections::HashMap; use std::sync::LazyLock; /// Global pricing data, loaded once from embedded JSON at first access. @@ -126,6 +127,8 @@ pub struct ModelPricingEntry { /// The full model pricing table. Loaded at startup from embedded JSON or `MODEL_PRICING_FILE`. pub struct ModelPricing { entries: Vec, + exact_index: HashMap, + prefix_indexes: Vec, } impl ModelPricing { @@ -157,7 +160,42 @@ impl ModelPricing { }; let entries: Vec = serde_json::from_str(&json).expect("invalid model_pricing.json"); - Self { entries } + Self::from_entries(entries) + } + + fn from_entries(entries: Vec) -> Self { + let mut exact_index = HashMap::with_capacity(entries.len()); + for (index, entry) in entries.iter().enumerate() { + exact_index + .entry(entry.model_pattern.clone()) + .or_insert(index); + } + + let mut prefix_indexes: Vec = (0..entries.len()).collect(); + prefix_indexes.sort_by(|left, right| { + entries[*right] + .model_pattern + .len() + .cmp(&entries[*left].model_pattern.len()) + .then_with(|| left.cmp(right)) + }); + + Self { + entries, + exact_index, + prefix_indexes, + } + } + + fn entry_for_model(&self, model: &str) -> Option<&ModelPricingEntry> { + if let Some(index) = self.exact_index.get(model) { + return self.entries.get(*index); + } + + self.prefix_indexes + .iter() + .map(|index| &self.entries[*index]) + .find(|entry| model.starts_with(&entry.model_pattern)) } /// Return (input_cost_per_token, output_cost_per_token) for a model, or None if unknown. @@ -165,18 +203,8 @@ impl ModelPricing { /// Same lookup order as cost_for_usage (exact then longest-prefix) but does not log /// on miss, so it is safe to call during routing decisions. pub fn price_for_model(&self, model: &str) -> Option<(f64, f64)> { - if let Some(entry) = self.entries.iter().find(|e| e.model_pattern == model) { - return Some((entry.input_cost_per_token, entry.output_cost_per_token)); - } - let mut best: Option<&ModelPricingEntry> = None; - let mut best_len: usize = 0; - for entry in &self.entries { - if model.starts_with(&entry.model_pattern) && entry.model_pattern.len() > best_len { - best = Some(entry); - best_len = entry.model_pattern.len(); - } - } - best.map(|e| (e.input_cost_per_token, e.output_cost_per_token)) + self.entry_for_model(model) + .map(|entry| (entry.input_cost_per_token, entry.output_cost_per_token)) } /// Calculate cost for a usage record. @@ -184,28 +212,11 @@ impl ModelPricing { /// Matching strategy: exact match first, then longest prefix match. /// Returns 0.0 with a warning log if no match found. pub fn cost_for_usage(&self, model: &str, input_tokens: u64, output_tokens: u64) -> f64 { - // 1. Try exact match - if let Some(entry) = self.entries.iter().find(|e| e.model_pattern == model) { + if let Some(entry) = self.entry_for_model(model) { return entry.input_cost_per_token * input_tokens as f64 + entry.output_cost_per_token * output_tokens as f64; } - // 2. Try longest prefix match (e.g., "gpt-4o-2024-05-13" matches "gpt-4o") - let mut best: Option<&ModelPricingEntry> = None; - let mut best_len: usize = 0; - for entry in &self.entries { - if model.starts_with(&entry.model_pattern) && entry.model_pattern.len() > best_len { - best = Some(entry); - best_len = entry.model_pattern.len(); - } - } - - if let Some(entry) = best { - return entry.input_cost_per_token * input_tokens as f64 - + entry.output_cost_per_token * output_tokens as f64; - } - - // 3. No match tracing::error!( model = model, "BILLING LEAK: no pricing entry found for model, cost set to 0.0" diff --git a/crates/proxy/src/cost/tests.rs b/crates/proxy/src/cost/tests.rs index b9c20ad..1b389b1 100644 --- a/crates/proxy/src/cost/tests.rs +++ b/crates/proxy/src/cost/tests.rs @@ -1,28 +1,26 @@ use super::*; fn test_pricing() -> ModelPricing { - ModelPricing { - entries: vec![ - ModelPricingEntry { - model_pattern: "gpt-4o".to_string(), - input_cost_per_token: 0.0000025, - output_cost_per_token: 0.00001, - provider: "openai".to_string(), - }, - ModelPricingEntry { - model_pattern: "gpt-4o-mini".to_string(), - input_cost_per_token: 0.00000015, - output_cost_per_token: 0.0000006, - provider: "openai".to_string(), - }, - ModelPricingEntry { - model_pattern: "gemini-2.5-pro".to_string(), - input_cost_per_token: 0.00000125, - output_cost_per_token: 0.00001, - provider: "google".to_string(), - }, - ], - } + ModelPricing::from_entries(vec![ + ModelPricingEntry { + model_pattern: "gpt-4o".to_string(), + input_cost_per_token: 0.0000025, + output_cost_per_token: 0.00001, + provider: "openai".to_string(), + }, + ModelPricingEntry { + model_pattern: "gpt-4o-mini".to_string(), + input_cost_per_token: 0.00000015, + output_cost_per_token: 0.0000006, + provider: "openai".to_string(), + }, + ModelPricingEntry { + model_pattern: "gemini-2.5-pro".to_string(), + input_cost_per_token: 0.00000125, + output_cost_per_token: 0.00001, + provider: "google".to_string(), + }, + ]) } #[test] diff --git a/crates/proxy/src/runtime/stream.rs b/crates/proxy/src/runtime/stream.rs index f19a9a2..f54cab5 100644 --- a/crates/proxy/src/runtime/stream.rs +++ b/crates/proxy/src/runtime/stream.rs @@ -1,4 +1,3 @@ -use bytes::BytesMut; use futures::StreamExt; use std::sync::Arc; use std::time::Instant; @@ -6,7 +5,7 @@ use tokio::sync::mpsc; use tokio_stream::wrappers::ReceiverStream; use super::{ChatCompletionChunkStream, ChatCompletionError}; -use crate::backend::MAX_SSE_BUFFER_SIZE; +use crate::backend::SseFrameBuffer; use anyllm_translate::{anthropic, mapping, openai, ReverseStreamingTranslator}; pub(crate) struct DeploymentLatencyGuard { @@ -48,8 +47,7 @@ pub(crate) fn openai_chunk_stream( tokio::spawn(async move { let fut = async { let mut byte_stream = response.bytes_stream(); - let mut buffer = BytesMut::new(); - let mut search_from: usize = 0; + let mut buffer = SseFrameBuffer::new(); while let Some(chunk_result) = byte_stream.next().await { let bytes = match chunk_result { @@ -61,16 +59,16 @@ pub(crate) fn openai_chunk_stream( } }; - if buffer.len() + bytes.len() > MAX_SSE_BUFFER_SIZE { - send_stream_error(&tx, ChatCompletionError::StreamBufferOverflow).await; - return; - } - buffer.extend_from_slice(&bytes); + let frames = match buffer.push(&bytes) { + Ok(frames) => frames, + Err(_) => { + send_stream_error(&tx, ChatCompletionError::StreamBufferOverflow).await; + return; + } + }; - while let Some((pos, delim_len)) = - anyllm_client::find_double_newline(&buffer, search_from) - { - let frame = match std::str::from_utf8(&buffer[..pos]) { + for frame in frames { + let frame = match std::str::from_utf8(&frame) { Ok(frame) => frame, Err(e) => { send_stream_error(&tx, ChatCompletionError::StreamParse(e.to_string())) @@ -109,11 +107,7 @@ pub(crate) fn openai_chunk_stream( } } } - - let _ = buffer.split_to(pos + delim_len); - search_from = 0; } - search_from = buffer.len().saturating_sub(3); } }; @@ -149,8 +143,7 @@ pub(crate) fn responses_chunk_stream( model, ); let mut byte_stream = response.bytes_stream(); - let mut buffer = BytesMut::new(); - let mut search_from: usize = 0; + let mut buffer = SseFrameBuffer::new(); while let Some(chunk_result) = byte_stream.next().await { let bytes = match chunk_result { @@ -162,16 +155,16 @@ pub(crate) fn responses_chunk_stream( } }; - if buffer.len() + bytes.len() > MAX_SSE_BUFFER_SIZE { - send_stream_error(&tx, ChatCompletionError::StreamBufferOverflow).await; - return; - } - buffer.extend_from_slice(&bytes); + let frames = match buffer.push(&bytes) { + Ok(frames) => frames, + Err(_) => { + send_stream_error(&tx, ChatCompletionError::StreamBufferOverflow).await; + return; + } + }; - while let Some((pos, delim_len)) = - anyllm_client::find_double_newline(&buffer, search_from) - { - let frame = match std::str::from_utf8(&buffer[..pos]) { + for frame in frames { + let frame = match std::str::from_utf8(&frame) { Ok(frame) => frame, Err(e) => { send_stream_error(&tx, ChatCompletionError::StreamParse(e.to_string())) @@ -212,11 +205,7 @@ pub(crate) fn responses_chunk_stream( } } } - - let _ = buffer.split_to(pos + delim_len); - search_from = 0; } - search_from = buffer.len().saturating_sub(3); } let final_events = responses_translator.finish(); diff --git a/crates/proxy/src/server/chat_completions/stream/anthropic.rs b/crates/proxy/src/server/chat_completions/stream/anthropic.rs index c55b258..d5ecd97 100644 --- a/crates/proxy/src/server/chat_completions/stream/anthropic.rs +++ b/crates/proxy/src/server/chat_completions/stream/anthropic.rs @@ -1,4 +1,4 @@ -use crate::backend::{find_double_newline, BackendError, MAX_SSE_BUFFER_SIZE}; +use crate::backend::{BackendError, SseFrameBuffer}; use crate::server::routes::{inject_degradation_header, log_request, set_backend_error_kind}; use crate::server::state::AppState; use crate::server::streaming::{AnthropicStreamUsage, StreamOutcome}; @@ -7,7 +7,6 @@ use axum::{ http::StatusCode, response::{IntoResponse, Response}, }; -use bytes::BytesMut; use futures::StreamExt; use crate::server::chat_completions::extensions::serialize_anthropic_upstream_request; @@ -94,8 +93,7 @@ pub(super) async fn anthropic_chat_completions_stream( ); let mut usage = AnthropicStreamUsage::default(); let mut byte_stream = response.bytes_stream(); - let mut buffer = BytesMut::new(); - let mut search_from: usize = 0; + let mut buffer = SseFrameBuffer::new(); let mut emitted_done = false; let stream_loop = async { @@ -109,18 +107,20 @@ pub(super) async fn anthropic_chat_completions_stream( } }; - if buffer.len() + bytes.len() > MAX_SSE_BUFFER_SIZE { - tracing::error!( - buffer_len = buffer.len(), - "Anthropic chat-completions SSE buffer exceeded maximum size" - ); - metrics.record_error(); - return StreamOutcome::UpstreamError; - } - buffer.extend_from_slice(&bytes); + let frames = match buffer.push(&bytes) { + Ok(frames) => frames, + Err(e) => { + tracing::error!( + error = %e, + "Anthropic chat-completions SSE buffer exceeded maximum size" + ); + metrics.record_error(); + return StreamOutcome::UpstreamError; + } + }; - while let Some((pos, delim_len)) = find_double_newline(&buffer, search_from) { - if let Ok(frame_str) = std::str::from_utf8(&buffer[..pos]) { + for frame in frames { + if let Ok(frame_str) = std::str::from_utf8(&frame) { for line in frame_str.lines() { let line = line.trim(); let Some(json_str) = line.strip_prefix("data: ") else { @@ -154,10 +154,7 @@ pub(super) async fn anthropic_chat_completions_stream( } } } - let _ = buffer.split_to(pos + delim_len); - search_from = 0; } - search_from = buffer.len().saturating_sub(3); } StreamOutcome::Completed }; diff --git a/crates/proxy/src/server/chat_completions/stream/generic/stream_loop.rs b/crates/proxy/src/server/chat_completions/stream/generic/stream_loop.rs index fbee757..f92e9b0 100644 --- a/crates/proxy/src/server/chat_completions/stream/generic/stream_loop.rs +++ b/crates/proxy/src/server/chat_completions/stream/generic/stream_loop.rs @@ -1,4 +1,4 @@ -use crate::backend::{find_double_newline, BackendClient, BackendError, MAX_SSE_BUFFER_SIZE}; +use crate::backend::{BackendClient, BackendError, SseFrameBuffer}; use crate::server::routes::{inject_degradation_header, log_request, set_backend_error_kind}; use crate::server::state::AppState; use anyllm_translate::{anthropic, mapping, openai, ReverseStreamingTranslator}; @@ -6,7 +6,6 @@ use axum::{ http::StatusCode, response::{IntoResponse, Response}, }; -use bytes::BytesMut; use futures::StreamExt; use crate::server::chat_completions::helpers::{ @@ -114,8 +113,7 @@ pub async fn generic_chat_completions_stream( mapping::streaming_map::StreamingTranslator::new(model_for_translator.clone()); let mut byte_stream = resp.bytes_stream(); - let mut buffer = BytesMut::new(); - let mut search_from: usize = 0; + let mut buffer = SseFrameBuffer::new(); let mut timed_out = false; // Accumulate tool call fragments for collect-then-execute. // Each entry: (id, function_name, arguments_json). @@ -132,18 +130,18 @@ pub async fn generic_chat_completions_stream( return; } }; - buffer.extend_from_slice(&bytes); + let frames = match buffer.push(&bytes) { + Ok(frames) => frames, + Err(e) => { + tracing::error!(error = %e, "SSE buffer exceeded maximum size"); + metrics.record_error(); + metrics.record_stream_failed(); + return; + } + }; - if buffer.len() > MAX_SSE_BUFFER_SIZE { - tracing::error!("SSE buffer exceeded maximum size"); - metrics.record_error(); - metrics.record_stream_failed(); - return; - } - - while let Some((pos, delim_len)) = find_double_newline(&buffer, search_from) - { - if let Ok(frame_str) = std::str::from_utf8(&buffer[..pos]) { + for frame in frames { + if let Ok(frame_str) = std::str::from_utf8(&frame) { for line in frame_str.lines() { let line = line.trim(); if let Some(json_str) = line.strip_prefix("data: ") { @@ -216,10 +214,7 @@ pub async fn generic_chat_completions_stream( } } } - let _ = buffer.split_to(pos + delim_len); - search_from = 0; } - search_from = buffer.len().saturating_sub(3); } }; diff --git a/crates/proxy/src/server/chat_completions/stream/generic/tool_loop.rs b/crates/proxy/src/server/chat_completions/stream/generic/tool_loop.rs index 94287e1..3850496 100644 --- a/crates/proxy/src/server/chat_completions/stream/generic/tool_loop.rs +++ b/crates/proxy/src/server/chat_completions/stream/generic/tool_loop.rs @@ -1,9 +1,8 @@ -use crate::backend::find_double_newline; use crate::backend::openai_client::OpenAIClient; +use crate::backend::SseFrameBuffer; use crate::server::middleware::VirtualKeyContext; use crate::server::state::ToolEngineState; use anyllm_translate::{anthropic, mapping, openai, ReverseStreamingTranslator}; -use bytes::BytesMut; use futures::StreamExt; #[allow(clippy::too_many_arguments)] @@ -116,8 +115,7 @@ pub(super) async fn run_tool_loop_for_stream( { Ok((follow_resp, _follow_rate_limits)) => { let mut follow_byte_stream = follow_resp.bytes_stream(); - let mut follow_buffer = BytesMut::new(); - let mut follow_search_from: usize = 0; + let mut follow_buffer = SseFrameBuffer::new(); while let Some(chunk_result) = follow_byte_stream.next().await { let bytes = match chunk_result { @@ -127,12 +125,16 @@ pub(super) async fn run_tool_loop_for_stream( break; } }; - follow_buffer.extend_from_slice(&bytes); + let frames = match follow_buffer.push(&bytes) { + Ok(frames) => frames, + Err(e) => { + tracing::error!(error = %e, "follow-up SSE buffer exceeded maximum size"); + break; + } + }; - while let Some((pos, delim_len)) = - find_double_newline(&follow_buffer, follow_search_from) - { - if let Ok(frame_str) = std::str::from_utf8(&follow_buffer[..pos]) { + for frame in frames { + if let Ok(frame_str) = std::str::from_utf8(&frame) { for line in frame_str.lines() { let line = line.trim(); if let Some(json_str) = line.strip_prefix("data: ") { @@ -198,10 +200,7 @@ pub(super) async fn run_tool_loop_for_stream( } } } - let _ = follow_buffer.split_to(pos + delim_len); - follow_search_from = 0; } - follow_search_from = follow_buffer.len().saturating_sub(3); } // Emit finish events for the follow-up stream. diff --git a/crates/proxy/src/server/gemini_input/stream.rs b/crates/proxy/src/server/gemini_input/stream.rs index 9fe14cf..255ec98 100644 --- a/crates/proxy/src/server/gemini_input/stream.rs +++ b/crates/proxy/src/server/gemini_input/stream.rs @@ -1,4 +1,4 @@ -use crate::backend::{find_double_newline, BackendClient, BackendError, MAX_SSE_BUFFER_SIZE}; +use crate::backend::{BackendClient, BackendError, SseFrameBuffer}; use crate::server::routes::{ backend_error_to_response, log_request, record_virtual_key_usage, set_backend_error_kind, RequestCtx, @@ -12,7 +12,6 @@ use axum::response::{ sse::{Event, KeepAlive, Sse}, IntoResponse, Response, }; -use bytes::BytesMut; use futures::StreamExt; use tokio::sync::mpsc; use tokio_stream::wrappers::ReceiverStream; @@ -74,9 +73,8 @@ pub(super) async fn gemini_stream( } }; - let mut buffer = BytesMut::new(); + let mut buffer = SseFrameBuffer::new(); let mut translator = streaming_map::StreamingTranslator::new(model.clone()); - let mut search_from: usize = 0; let mut byte_stream = response.bytes_stream(); let mut outcome = StreamOutcome::Completed; @@ -93,16 +91,21 @@ pub(super) async fn gemini_stream( } }; - if buffer.len() + bytes.len() > MAX_SSE_BUFFER_SIZE { - tracing::error!("SSE buffer exceeded max, aborting gemini input stream"); - metrics.record_error(); - outcome = StreamOutcome::UpstreamError; - break; - } - buffer.extend_from_slice(&bytes); + let frames = match buffer.push(&bytes) { + Ok(frames) => frames, + Err(e) => { + tracing::error!( + error = %e, + "SSE buffer exceeded max, aborting gemini input stream" + ); + metrics.record_error(); + outcome = StreamOutcome::UpstreamError; + break; + } + }; - while let Some((pos, delim_len)) = find_double_newline(&buffer, search_from) { - if let Ok(frame_str) = std::str::from_utf8(&buffer[..pos]) { + for frame in frames { + if let Ok(frame_str) = std::str::from_utf8(&frame) { for line in frame_str.lines() { let line = line.trim(); if let Some(json_str) = line.strip_prefix("data: ") { @@ -136,10 +139,7 @@ pub(super) async fn gemini_stream( } } } - let _ = buffer.split_to(pos + delim_len); - search_from = 0; } - search_from = buffer.len().saturating_sub(3); } if matches!(outcome, StreamOutcome::Completed) && !done { let _ = translator.finish(); diff --git a/crates/proxy/src/server/routes/handlers.rs b/crates/proxy/src/server/routes/handlers.rs index 52235e3..e3b6138 100644 --- a/crates/proxy/src/server/routes/handlers.rs +++ b/crates/proxy/src/server/routes/handlers.rs @@ -3,6 +3,21 @@ use crate::server::state::AppState; use anyllm_providers::ProviderCatalog; use axum::extract::State; use axum::response::{IntoResponse, Json, Response}; +use std::collections::{HashMap, HashSet}; +use std::sync::{Arc, LazyLock, Mutex, Weak}; + +struct CachedAnthropicCatalogRows { + catalog: Weak, + rows: Arc, +} + +struct AnthropicCatalogRows { + rows: Arc<[serde_json::Value]>, + ids: Arc>, +} + +static ANTHROPIC_CATALOG_ROWS_CACHE: LazyLock>> = + LazyLock::new(|| Mutex::new(HashMap::new())); pub(crate) fn anthropic_catalog_model_rows(catalog: &ProviderCatalog) -> Vec { catalog @@ -20,6 +35,40 @@ pub(crate) fn anthropic_catalog_model_rows(catalog: &ProviderCatalog) -> Vec, +) -> Arc { + let key = Arc::as_ptr(catalog) as usize; + let mut cache = ANTHROPIC_CATALOG_ROWS_CACHE + .lock() + .unwrap_or_else(|e| e.into_inner()); + if let Some(entry) = cache.get(&key) { + if let Some(cached_catalog) = entry.catalog.upgrade() { + if Arc::ptr_eq(&cached_catalog, catalog) { + return entry.rows.clone(); + } + } + } + + let rows = anthropic_catalog_model_rows(catalog); + let ids = rows + .iter() + .filter_map(|model| model["id"].as_str().map(str::to_string)) + .collect(); + let rows = Arc::new(AnthropicCatalogRows { + rows: Arc::from(rows), + ids: Arc::new(ids), + }); + cache.insert( + key, + CachedAnthropicCatalogRows { + catalog: Arc::downgrade(catalog), + rows: rows.clone(), + }, + ); + rows +} + fn claude_display_name(model_id: &str) -> String { let name = model_id .strip_prefix("claude-") @@ -46,17 +95,14 @@ fn claude_display_name(model_id: &str) -> String { /// GET /v1/models -- returns catalog Claude models merged with model_list entries. pub async fn models(State(state): State) -> Json { - let mut data = anthropic_catalog_model_rows(&state.provider_catalog); + let cached_rows = cached_anthropic_catalog_model_rows(&state.provider_catalog); + let mut data = cached_rows.rows.iter().cloned().collect::>(); // Merge models from the model router (LiteLLM model_list config). if let Some(ref router_lock) = state.model_router { let router = router_lock.read().unwrap_or_else(|e| e.into_inner()); - let static_ids: std::collections::HashSet = data - .iter() - .filter_map(|m| m["id"].as_str().map(|s| s.to_string())) - .collect(); for model_name in router.known_models() { - if !static_ids.contains(model_name) { + if !cached_rows.ids.contains(model_name) { data.push(serde_json::json!({ "id": model_name, "object": "model", diff --git a/crates/proxy/src/server/streaming.rs b/crates/proxy/src/server/streaming.rs index a58b9ed..047ece3 100644 --- a/crates/proxy/src/server/streaming.rs +++ b/crates/proxy/src/server/streaming.rs @@ -1,6 +1,6 @@ // SSE streaming infrastructure and the messages_stream handler. -use crate::backend::{find_double_newline, BackendClient, RateLimitHeaders, MAX_SSE_BUFFER_SIZE}; +use crate::backend::{find_double_newline, BackendClient, RateLimitHeaders, SseFrameBuffer}; use crate::metrics::Metrics; use anyllm_translate::{anthropic, mapping, openai}; use axum::response::sse::{Event, KeepAlive, Sse}; @@ -172,15 +172,12 @@ where { use futures::StreamExt; let mut stream = response.bytes_stream(); - // BytesMut (not String) because TCP chunks may split mid-UTF-8 character. + // Buffer bytes (not String) because TCP chunks may split mid-UTF-8 character. // String::from_utf8_lossy would permanently replace partial trailing bytes // with U+FFFD, corrupting the JSON payload. - let mut buffer = BytesMut::new(); + let mut buffer = SseFrameBuffer::new(); // Reuse a single events buffer across all frames to avoid per-frame allocation let mut frame_events: Vec = Vec::new(); - // Track where to start the next delimiter search so we don't rescan - // already-inspected bytes when a large SSE event spans many TCP chunks. - let mut search_from: usize = 0; while let Some(chunk_result) = stream.next().await { let bytes = match chunk_result { @@ -191,24 +188,21 @@ where return StreamOutcome::UpstreamError; } }; - // Guard against unbounded buffer growth from a misbehaving backend. - // Check before appending so a single oversized chunk can't exceed the limit. - if buffer.len() + bytes.len() > MAX_SSE_BUFFER_SIZE { - tracing::error!( - buffer_len = buffer.len(), - "SSE buffer exceeded maximum size, aborting stream" - ); - metrics.record_error(); - return StreamOutcome::UpstreamError; - } - buffer.extend_from_slice(&bytes); + let frames = match buffer.push(&bytes) { + Ok(frames) => frames, + Err(e) => { + tracing::error!(error = %e, "SSE buffer exceeded maximum size, aborting stream"); + metrics.record_error(); + return StreamOutcome::UpstreamError; + } + }; - while let Some((pos, delim_len)) = find_double_newline(&buffer, search_from) { + for frame in frames { frame_events.clear(); // Convert the complete frame bytes to UTF-8. A frame ending at // a double-newline boundary should always be valid UTF-8; if not, // skip the malformed frame rather than injecting replacement chars. - match std::str::from_utf8(&buffer[..pos]) { + match std::str::from_utf8(&frame) { Ok(frame_str) => { for line in frame_str.lines() { let line = line.trim(); @@ -223,19 +217,12 @@ where tracing::warn!("skipping non-UTF-8 SSE frame: {e}"); } } - let _ = buffer.split_to(pos + delim_len); - // split_to shifted the buffer; restart search at the beginning - search_from = 0; if !send_events(tx, &frame_events).await { tracing::debug!("client disconnected during stream"); return StreamOutcome::ClientDisconnected; } } - // Next chunk: resume scanning 3 bytes back from the end. The 4-byte - // delimiter \r\n\r\n could straddle the chunk boundary (e.g., \r\n at - // end of this chunk, \r\n at start of the next). - search_from = buffer.len().saturating_sub(3); } StreamOutcome::Completed diff --git a/docs/providers/bedrock.md b/docs/providers/bedrock.md index ed85500..63e6726 100644 --- a/docs/providers/bedrock.md +++ b/docs/providers/bedrock.md @@ -50,6 +50,8 @@ model_list: aws_region_name: us-east-1 ``` +LiteLLM YAML compatibility is limited to `aws_access_key_id`, `aws_secret_access_key`, and `aws_region_name`, with access key and secret falling back to `AWS_ACCESS_KEY_ID` and `AWS_SECRET_ACCESS_KEY`. Bedrock bearer API keys, profiles, role assumption, web identity, custom runtime endpoints, and `aws_session_token` in LiteLLM YAML are not implemented. + ## Usage Examples ### Anthropic Messages API diff --git a/docs/providers/vertex_ai.md b/docs/providers/vertex_ai.md index 02ab21f..c82c90f 100644 --- a/docs/providers/vertex_ai.md +++ b/docs/providers/vertex_ai.md @@ -12,11 +12,10 @@ Google Vertex AI — enterprise Gemini and third-party models via GCP. |---|---|---| | `VERTEX_PROJECT` | Yes | GCP project ID (e.g. `my-project-123`) | | `VERTEX_REGION` | Yes | GCP region (e.g. `us-central1`) | -| `GOOGLE_APPLICATION_CREDENTIALS` | Yes (or alt) | Path to service account JSON key file | -| `VERTEX_API_KEY` | Yes (or alt) | API key if not using a service account | -| `GOOGLE_ACCESS_TOKEN` | No | Short-lived bearer token (overrides key auth) | +| `VERTEX_API_KEY` | Yes (or alt) | Google API key | +| `GOOGLE_ACCESS_TOKEN` | Yes (or alt) | Short-lived bearer token | -Provide either `GOOGLE_APPLICATION_CREDENTIALS` (service account) or `VERTEX_API_KEY`, not both. +Provide either `VERTEX_API_KEY` or `GOOGLE_ACCESS_TOKEN`. Service-account JSON loading from `GOOGLE_APPLICATION_CREDENTIALS` is not implemented; mint a token externally and pass it via `GOOGLE_ACCESS_TOKEN`. ## Quick Start @@ -26,15 +25,14 @@ Provide either `GOOGLE_APPLICATION_CREDENTIALS` (service account) or `VERTEX_API BACKEND=vertex_ai \ VERTEX_PROJECT=my-project-123 \ VERTEX_REGION=us-central1 \ - GOOGLE_APPLICATION_CREDENTIALS=/path/to/sa.json \ + VERTEX_API_KEY=AIza... \ cargo run -p anyllm_proxy # or with Docker: docker run \ -e BACKEND=vertex_ai \ -e VERTEX_PROJECT=my-project-123 \ -e VERTEX_REGION=us-central1 \ - -e GOOGLE_APPLICATION_CREDENTIALS=/run/secrets/sa.json \ - -v /path/to/sa.json:/run/secrets/sa.json:ro \ + -e VERTEX_API_KEY=AIza... \ -e PROXY_OPEN_RELAY=true \ -p 3000:3000 \ followthewhit3rabbit/anyllm-proxy @@ -47,11 +45,13 @@ model_list: - model_name: gemini-2.5-pro litellm_params: model: vertex_ai/gemini-2.5-pro + api_key: os.environ/VERTEX_API_KEY vertex_project: my-project-123 vertex_location: us-central1 - model_name: claude-3-5-sonnet-vertex litellm_params: model: vertex_ai/claude-3-5-sonnet@20241022 + api_key: os.environ/GOOGLE_ACCESS_TOKEN vertex_project: my-project-123 vertex_location: us-east5 ``` @@ -106,8 +106,8 @@ curl http://localhost:3000/v1/chat/completions \ ## Notes -- The base URL is constructed per request: `https://{VERTEX_REGION}-aiplatform.googleapis.com/v1/projects/{VERTEX_PROJECT}/locations/{VERTEX_REGION}/publishers/google/models/{model}`. +- The OpenAI-compatible base URL is constructed from project and region: `https://{VERTEX_REGION}-aiplatform.googleapis.com/v1/projects/{VERTEX_PROJECT}/locations/{VERTEX_REGION}/endpoints/openapi`. - Vertex AI serves the same Gemini model IDs as Google AI Studio but requires a GCP project with the Vertex AI API enabled (`gcloud services enable aiplatform.googleapis.com`). - Claude models (Anthropic Model Garden) use region-specific availability. `us-east5` is the primary region for Claude on Vertex; check the GCP console for current availability. -- Service account must have the `roles/aiplatform.user` IAM role. +- If you mint `GOOGLE_ACCESS_TOKEN` from a service account, that account must have the `roles/aiplatform.user` IAM role. - No static model list is maintained in the proxy. Pass the model ID directly as it appears in the Vertex API.