mirror of
https://github.com/whit3rabbit/anyllm-proxy.git
synced 2026-09-22 00:00:50 +00:00
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:
co-authored by
Claude Opus 4.8
parent
af7fe570c7
commit
ee6a39d5d4
@@ -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
|
||||
|
||||
@@ -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>
|
||||
|
||||
|
||||
@@ -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
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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,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": {
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
Vendored
+148
-44
@@ -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));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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#"
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user