fix: streaming/cost/cache correctness, litellm compat, provider docs

- client/proxy SSE + streaming loop cleanups (sse.rs, streaming.rs,
  runtime/stream.rs, chat_completions stream/anthropic/generic, gemini_input)
- cost + cache module fixes with added tests
- provider catalog/registry adjustments
- litellm YAML config support + tests; single-config tweaks
- docs: scope Vertex auth to VERTEX_API_KEY/GOOGLE_ACCESS_TOKEN,
  Bedrock LiteLLM YAML limits, Docker Hub image/tag docs in README
- CI: docker workflow multi-arch tag handling

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
whit3rabbit
2026-06-20 12:24:07 -05:00
co-authored by Claude Opus 4.8
parent af7fe570c7
commit ee6a39d5d4
24 changed files with 803 additions and 318 deletions
+32 -3
View File
@@ -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
+8
View File
@@ -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-<short-sha>` for the release commit
- `latest` for stable `vX.Y.Z` releases only, prerelease tags do not update it
<details>
<summary>Smoke tests (no real API key needed)</summary>
+1 -1
View File
@@ -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
+85 -16
View File
@@ -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<Vec<Bytes>, 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<T> = 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");
}
}
+11 -17
View File
@@ -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<Result<StreamEvent, ClientError>>,
) {
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) {
+33 -9
View File
@@ -58,6 +58,8 @@ pub struct ProviderCatalog {
providers: BTreeMap<String, OwnedProviderDef>,
advertised_provider_ids: BTreeSet<String>,
models_by_provider: BTreeMap<String, Vec<OwnedModelDef>>,
provider_ids_by_litellm_prefix: BTreeMap<String, String>,
model_indexes_by_provider: BTreeMap<String, BTreeMap<String, usize>>,
}
#[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)]
+1 -1
View File
@@ -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": {
+65 -30
View File
@@ -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<HashMap<&'static str, &'static ProviderDef>> =
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<HashMap<&'static str, &'static ProviderDef>> =
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<HashMap<&'static str, &'static [ModelDef]>> =
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<Item = &'static ProviderDef> {
/// 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)]
+1 -1
View File
@@ -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;
+148 -44
View File
@@ -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<usize> {
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));
}
}
+49 -5
View File
@@ -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<u32>,
// Azure-specific
api_version: Option<String>,
// Vertex-specific
vertex_project: Option<String>,
vertex_location: Option<String>,
// Bedrock-specific
aws_access_key_id: Option<String>,
aws_secret_access_key: Option<String>,
@@ -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 {
+125
View File
@@ -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#"
+30 -1
View File
@@ -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");
}
}
}
+42 -31
View File
@@ -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<ModelPricingEntry>,
exact_index: HashMap<String, usize>,
prefix_indexes: Vec<usize>,
}
impl ModelPricing {
@@ -157,7 +160,42 @@ impl ModelPricing {
};
let entries: Vec<ModelPricingEntry> =
serde_json::from_str(&json).expect("invalid model_pricing.json");
Self { entries }
Self::from_entries(entries)
}
fn from_entries(entries: Vec<ModelPricingEntry>) -> 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<usize> = (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"
+20 -22
View File
@@ -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]
+21 -32
View File
@@ -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();
@@ -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
};
@@ -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);
}
};
@@ -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.
+16 -16
View File
@@ -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();
+52 -6
View File
@@ -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<ProviderCatalog>,
rows: Arc<AnthropicCatalogRows>,
}
struct AnthropicCatalogRows {
rows: Arc<[serde_json::Value]>,
ids: Arc<HashSet<String>>,
}
static ANTHROPIC_CATALOG_ROWS_CACHE: LazyLock<Mutex<HashMap<usize, CachedAnthropicCatalogRows>>> =
LazyLock::new(|| Mutex::new(HashMap::new()));
pub(crate) fn anthropic_catalog_model_rows(catalog: &ProviderCatalog) -> Vec<serde_json::Value> {
catalog
@@ -20,6 +35,40 @@ pub(crate) fn anthropic_catalog_model_rows(catalog: &ProviderCatalog) -> Vec<ser
.collect()
}
fn cached_anthropic_catalog_model_rows(
catalog: &Arc<ProviderCatalog>,
) -> Arc<AnthropicCatalogRows> {
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<AppState>) -> Json<serde_json::Value> {
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::<Vec<_>>();
// 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<String> = 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",
+13 -26
View File
@@ -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<anthropic::StreamEvent> = 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
+2
View File
@@ -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
+9 -9
View File
@@ -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.