cachellm-proxy 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- cachellm/__init__.py +13 -0
- cachellm/__main__.py +6 -0
- cachellm/api/__init__.py +5 -0
- cachellm/api/app.py +143 -0
- cachellm/api/auth.py +34 -0
- cachellm/api/deps.py +106 -0
- cachellm/api/routes_admin.py +265 -0
- cachellm/api/routes_chat.py +385 -0
- cachellm/api/sse.py +98 -0
- cachellm/cache/__init__.py +3 -0
- cachellm/cache/analytics.py +150 -0
- cachellm/cache/coalesce.py +63 -0
- cachellm/cache/entry.py +92 -0
- cachellm/cache/exact_store.py +33 -0
- cachellm/cache/keys.py +124 -0
- cachellm/cache/policy.py +134 -0
- cachellm/cache/redis_client.py +22 -0
- cachellm/cache/service.py +332 -0
- cachellm/cache/vector_store.py +217 -0
- cachellm/cli.py +122 -0
- cachellm/embeddings/__init__.py +19 -0
- cachellm/embeddings/base.py +38 -0
- cachellm/embeddings/fastembed_backend.py +75 -0
- cachellm/embeddings/hash_backend.py +42 -0
- cachellm/errors.py +72 -0
- cachellm/logging_setup.py +56 -0
- cachellm/models.py +181 -0
- cachellm/observability/__init__.py +6 -0
- cachellm/observability/metrics.py +147 -0
- cachellm/observability/tracing.py +107 -0
- cachellm/pricing.py +108 -0
- cachellm/providers/__init__.py +7 -0
- cachellm/providers/base.py +84 -0
- cachellm/providers/bedrock.py +238 -0
- cachellm/providers/fake.py +56 -0
- cachellm/providers/openai_compat.py +131 -0
- cachellm/providers/registry.py +96 -0
- cachellm/py.typed +0 -0
- cachellm/settings.py +230 -0
- cachellm_proxy-0.1.0.dist-info/METADATA +550 -0
- cachellm_proxy-0.1.0.dist-info/RECORD +43 -0
- cachellm_proxy-0.1.0.dist-info/WHEEL +4 -0
- cachellm_proxy-0.1.0.dist-info/entry_points.txt +3 -0
|
@@ -0,0 +1,107 @@
|
|
|
1
|
+
"""Optional OpenTelemetry tracing.
|
|
2
|
+
|
|
3
|
+
OTel is the spine, not the destination. Instrument once and the same spans can
|
|
4
|
+
go to Langfuse (which reads them as LLM generations, with cost and prompts),
|
|
5
|
+
Grafana Tempo, Jaeger, or anything else that accepts OTLP. Nothing here is
|
|
6
|
+
required for the proxy to run: with tracing disabled every helper degrades to a
|
|
7
|
+
no-op context manager.
|
|
8
|
+
|
|
9
|
+
Span attributes follow the OpenTelemetry GenAI semantic conventions where they
|
|
10
|
+
exist (``gen_ai.*``), with CacheLLM's own facts under ``cachellm.*``.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import base64
|
|
16
|
+
import contextlib
|
|
17
|
+
import os
|
|
18
|
+
from collections.abc import Iterator
|
|
19
|
+
from typing import Any
|
|
20
|
+
|
|
21
|
+
import structlog
|
|
22
|
+
|
|
23
|
+
from cachellm.settings import Settings
|
|
24
|
+
|
|
25
|
+
log = structlog.get_logger(__name__)
|
|
26
|
+
|
|
27
|
+
_tracer: Any = None
|
|
28
|
+
_enabled = False
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def langfuse_otlp_headers(public_key: str, secret_key: str) -> str:
|
|
32
|
+
"""Langfuse authenticates OTLP with HTTP basic auth over the key pair."""
|
|
33
|
+
token = base64.b64encode(f"{public_key}:{secret_key}".encode()).decode()
|
|
34
|
+
return f"Authorization=Basic {token}"
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def setup_tracing(settings: Settings, app: Any = None) -> bool:
|
|
38
|
+
"""Configure the global tracer. Returns True when tracing is live."""
|
|
39
|
+
global _tracer, _enabled
|
|
40
|
+
if not settings.tracing_enabled:
|
|
41
|
+
return False
|
|
42
|
+
try:
|
|
43
|
+
from opentelemetry import trace
|
|
44
|
+
from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter
|
|
45
|
+
from opentelemetry.sdk.resources import Resource
|
|
46
|
+
from opentelemetry.sdk.trace import TracerProvider
|
|
47
|
+
from opentelemetry.sdk.trace.export import BatchSpanProcessor
|
|
48
|
+
except ImportError:
|
|
49
|
+
log.warning(
|
|
50
|
+
"tracing_unavailable",
|
|
51
|
+
hint="install the observability extra: uv sync --extra observability",
|
|
52
|
+
)
|
|
53
|
+
return False
|
|
54
|
+
|
|
55
|
+
endpoint = settings.otlp_endpoint or os.getenv("OTEL_EXPORTER_OTLP_ENDPOINT", "")
|
|
56
|
+
if not endpoint:
|
|
57
|
+
log.warning("tracing_disabled_no_endpoint")
|
|
58
|
+
return False
|
|
59
|
+
|
|
60
|
+
provider = TracerProvider(resource=Resource.create({"service.name": settings.service_name}))
|
|
61
|
+
provider.add_span_processor(BatchSpanProcessor(OTLPSpanExporter(endpoint=endpoint)))
|
|
62
|
+
trace.set_tracer_provider(provider)
|
|
63
|
+
_tracer = trace.get_tracer("cachellm")
|
|
64
|
+
_enabled = True
|
|
65
|
+
|
|
66
|
+
for module, instrument in (
|
|
67
|
+
("opentelemetry.instrumentation.redis", "RedisInstrumentor"),
|
|
68
|
+
("opentelemetry.instrumentation.httpx", "HTTPXClientInstrumentor"),
|
|
69
|
+
("opentelemetry.instrumentation.botocore", "BotocoreInstrumentor"),
|
|
70
|
+
):
|
|
71
|
+
with contextlib.suppress(Exception):
|
|
72
|
+
mod = __import__(module, fromlist=[instrument])
|
|
73
|
+
getattr(mod, instrument)().instrument()
|
|
74
|
+
|
|
75
|
+
if app is not None:
|
|
76
|
+
with contextlib.suppress(Exception):
|
|
77
|
+
from opentelemetry.instrumentation.fastapi import FastAPIInstrumentor
|
|
78
|
+
|
|
79
|
+
FastAPIInstrumentor.instrument_app(app, excluded_urls="healthz,readyz,metrics")
|
|
80
|
+
|
|
81
|
+
log.info("tracing_enabled", endpoint=endpoint, service=settings.service_name)
|
|
82
|
+
return True
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
@contextlib.contextmanager
|
|
86
|
+
def span(name: str, **attributes: Any) -> Iterator[Any]:
|
|
87
|
+
"""Start a span, or do nothing at all when tracing is off."""
|
|
88
|
+
if not _enabled or _tracer is None:
|
|
89
|
+
yield None
|
|
90
|
+
return
|
|
91
|
+
with _tracer.start_as_current_span(name) as current:
|
|
92
|
+
for key, value in attributes.items():
|
|
93
|
+
if value is not None:
|
|
94
|
+
current.set_attribute(key, value)
|
|
95
|
+
yield current
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def set_attributes(current: Any, **attributes: Any) -> None:
|
|
99
|
+
if current is None:
|
|
100
|
+
return
|
|
101
|
+
for key, value in attributes.items():
|
|
102
|
+
if value is not None:
|
|
103
|
+
current.set_attribute(key, value)
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def is_enabled() -> bool:
|
|
107
|
+
return _enabled
|
cachellm/pricing.py
ADDED
|
@@ -0,0 +1,108 @@
|
|
|
1
|
+
"""Token prices used to compute "money saved" for every cache hit.
|
|
2
|
+
|
|
3
|
+
Prices are USD per one million tokens and are *defaults*: providers change them,
|
|
4
|
+
so treat this table as a starting point and override it with
|
|
5
|
+
``CACHELLM_PRICING_FILE`` pointing at a JSON file of the same shape.
|
|
6
|
+
|
|
7
|
+
The saving reported by CacheLLM is a modelled number, not a bill. It is
|
|
8
|
+
"what these tokens would have cost at list price if we had called the provider",
|
|
9
|
+
which is exactly the number a team wants when deciding whether to deploy this.
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
from __future__ import annotations
|
|
13
|
+
|
|
14
|
+
import json
|
|
15
|
+
import os
|
|
16
|
+
from dataclasses import dataclass
|
|
17
|
+
|
|
18
|
+
__all__ = ["PRICES", "ModelPrice", "estimate_cost", "price_for"]
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
@dataclass(frozen=True)
|
|
22
|
+
class ModelPrice:
|
|
23
|
+
"""USD per 1M tokens."""
|
|
24
|
+
|
|
25
|
+
input_per_m: float
|
|
26
|
+
output_per_m: float
|
|
27
|
+
|
|
28
|
+
def cost(self, input_tokens: int, output_tokens: int) -> float:
|
|
29
|
+
return (input_tokens * self.input_per_m + output_tokens * self.output_per_m) / 1_000_000
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
# Defaults verified against public list prices; override per deployment.
|
|
33
|
+
PRICES: dict[str, ModelPrice] = {
|
|
34
|
+
# --- AWS Bedrock: Amazon Nova ---
|
|
35
|
+
"amazon.nova-micro-v1:0": ModelPrice(0.035, 0.14),
|
|
36
|
+
"amazon.nova-lite-v1:0": ModelPrice(0.06, 0.24),
|
|
37
|
+
"amazon.nova-pro-v1:0": ModelPrice(0.80, 3.20),
|
|
38
|
+
# --- AWS Bedrock: Anthropic ---
|
|
39
|
+
"anthropic.claude-3-haiku-20240307-v1:0": ModelPrice(0.25, 1.25),
|
|
40
|
+
"anthropic.claude-3-5-haiku-20241022-v1:0": ModelPrice(0.80, 4.00),
|
|
41
|
+
"anthropic.claude-3-sonnet-20240229-v1:0": ModelPrice(3.00, 15.00),
|
|
42
|
+
"anthropic.claude-3-5-sonnet-20240620-v1:0": ModelPrice(3.00, 15.00),
|
|
43
|
+
# --- AWS Bedrock: Meta / Mistral ---
|
|
44
|
+
"meta.llama3-8b-instruct-v1:0": ModelPrice(0.30, 0.60),
|
|
45
|
+
"meta.llama3-70b-instruct-v1:0": ModelPrice(2.65, 3.50),
|
|
46
|
+
"mistral.mistral-7b-instruct-v0:2": ModelPrice(0.15, 0.20),
|
|
47
|
+
"mistral.mixtral-8x7b-instruct-v0:1": ModelPrice(0.45, 0.70),
|
|
48
|
+
"mistral.mistral-large-2402-v1:0": ModelPrice(4.00, 12.00),
|
|
49
|
+
# --- OpenAI ---
|
|
50
|
+
"gpt-4o-mini": ModelPrice(0.15, 0.60),
|
|
51
|
+
"gpt-4o": ModelPrice(2.50, 10.00),
|
|
52
|
+
"gpt-4.1-mini": ModelPrice(0.40, 1.60),
|
|
53
|
+
# --- embeddings (input only) ---
|
|
54
|
+
"amazon.titan-embed-text-v2:0": ModelPrice(0.02, 0.0),
|
|
55
|
+
"text-embedding-3-small": ModelPrice(0.02, 0.0),
|
|
56
|
+
"text-embedding-3-large": ModelPrice(0.13, 0.0),
|
|
57
|
+
# --- test double ---
|
|
58
|
+
"fake-echo": ModelPrice(0.50, 1.50),
|
|
59
|
+
}
|
|
60
|
+
|
|
61
|
+
_FALLBACK = ModelPrice(0.15, 0.60)
|
|
62
|
+
_loaded_override = False
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def _load_override() -> None:
|
|
66
|
+
global _loaded_override
|
|
67
|
+
if _loaded_override:
|
|
68
|
+
return
|
|
69
|
+
_loaded_override = True
|
|
70
|
+
path = os.getenv("CACHELLM_PRICING_FILE", "").strip()
|
|
71
|
+
if not path or not os.path.exists(path):
|
|
72
|
+
return
|
|
73
|
+
with open(path, encoding="utf-8") as fh:
|
|
74
|
+
raw = json.load(fh)
|
|
75
|
+
for model, entry in raw.items():
|
|
76
|
+
PRICES[model] = ModelPrice(float(entry["input_per_m"]), float(entry["output_per_m"]))
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def normalise_model_id(model: str) -> str:
|
|
80
|
+
"""Strip routing and inference-profile decorations from a model id.
|
|
81
|
+
|
|
82
|
+
``bedrock/us.amazon.nova-lite-v1:0`` and ``amazon.nova-lite-v1:0`` are the
|
|
83
|
+
same billable model, and cross-region profiles just add a region prefix.
|
|
84
|
+
"""
|
|
85
|
+
m = model.strip()
|
|
86
|
+
if "/" in m:
|
|
87
|
+
m = m.split("/", 1)[1]
|
|
88
|
+
for region_prefix in ("us.", "eu.", "apac.", "ap."):
|
|
89
|
+
if m.startswith(region_prefix):
|
|
90
|
+
m = m[len(region_prefix) :]
|
|
91
|
+
break
|
|
92
|
+
return m
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def price_for(model: str) -> ModelPrice:
|
|
96
|
+
_load_override()
|
|
97
|
+
key = normalise_model_id(model)
|
|
98
|
+
if key in PRICES:
|
|
99
|
+
return PRICES[key]
|
|
100
|
+
# Prefix match so undated model revisions still price sensibly.
|
|
101
|
+
for known, price in PRICES.items():
|
|
102
|
+
if key.startswith(known.split("-2024")[0].split("-v1:")[0]):
|
|
103
|
+
return price
|
|
104
|
+
return _FALLBACK
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def estimate_cost(model: str, input_tokens: int, output_tokens: int) -> float:
|
|
108
|
+
return price_for(model).cost(input_tokens, output_tokens)
|
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from cachellm.providers.base import Provider, ProviderResult, StreamEvent
|
|
4
|
+
from cachellm.providers.fake import FakeProvider
|
|
5
|
+
from cachellm.providers.registry import ProviderRegistry
|
|
6
|
+
|
|
7
|
+
__all__ = ["FakeProvider", "Provider", "ProviderRegistry", "ProviderResult", "StreamEvent"]
|
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
"""Provider interface.
|
|
2
|
+
|
|
3
|
+
Everything the cache needs from an upstream model lives behind two methods.
|
|
4
|
+
Keeping this surface tiny is what makes the cache provider-agnostic: a cached
|
|
5
|
+
Bedrock answer never gets served to an OpenAI request (provider is part of the
|
|
6
|
+
namespace), but the caching logic itself is written once.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import abc
|
|
12
|
+
from collections.abc import AsyncIterator
|
|
13
|
+
from dataclasses import dataclass, field
|
|
14
|
+
from typing import Any
|
|
15
|
+
|
|
16
|
+
from cachellm.models import ChatCompletionRequest
|
|
17
|
+
|
|
18
|
+
# Upstream "finished cleanly" markers, normalised to the OpenAI vocabulary.
|
|
19
|
+
FINISH_REASON_MAP = {
|
|
20
|
+
"end_turn": "stop",
|
|
21
|
+
"stop_sequence": "stop",
|
|
22
|
+
"stop": "stop",
|
|
23
|
+
"max_tokens": "length",
|
|
24
|
+
"length": "length",
|
|
25
|
+
"content_filtered": "content_filter",
|
|
26
|
+
"content_filter": "content_filter",
|
|
27
|
+
"tool_use": "tool_calls",
|
|
28
|
+
"tool_calls": "tool_calls",
|
|
29
|
+
"guardrail_intervened": "content_filter",
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def normalise_finish_reason(raw: str | None) -> str:
|
|
34
|
+
if not raw:
|
|
35
|
+
return "stop"
|
|
36
|
+
return FINISH_REASON_MAP.get(str(raw).lower(), str(raw).lower())
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@dataclass
|
|
40
|
+
class ProviderResult:
|
|
41
|
+
text: str
|
|
42
|
+
model: str
|
|
43
|
+
prompt_tokens: int = 0
|
|
44
|
+
completion_tokens: int = 0
|
|
45
|
+
finish_reason: str = "stop"
|
|
46
|
+
tool_calls: list[dict[str, Any]] | None = None
|
|
47
|
+
raw: dict[str, Any] = field(default_factory=dict)
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
@dataclass
|
|
51
|
+
class StreamEvent:
|
|
52
|
+
"""One step of a streamed completion."""
|
|
53
|
+
|
|
54
|
+
delta: str = ""
|
|
55
|
+
finish_reason: str | None = None
|
|
56
|
+
prompt_tokens: int = 0
|
|
57
|
+
completion_tokens: int = 0
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
class Provider(abc.ABC):
|
|
61
|
+
name: str
|
|
62
|
+
|
|
63
|
+
@abc.abstractmethod
|
|
64
|
+
async def complete(self, request: ChatCompletionRequest) -> ProviderResult: ...
|
|
65
|
+
|
|
66
|
+
@abc.abstractmethod
|
|
67
|
+
def stream(self, request: ChatCompletionRequest) -> AsyncIterator[StreamEvent]: ...
|
|
68
|
+
|
|
69
|
+
def resolve_model(self, model: str) -> str:
|
|
70
|
+
"""Strip the routing prefix, if any, to get the upstream model id."""
|
|
71
|
+
return model.split("/", 1)[1] if "/" in model else model
|
|
72
|
+
|
|
73
|
+
async def close(self) -> None: # pragma: no cover - default no-op
|
|
74
|
+
return None
|
|
75
|
+
|
|
76
|
+
@staticmethod
|
|
77
|
+
def approx_tokens(text: str) -> int:
|
|
78
|
+
"""Rough token count for providers that do not report usage.
|
|
79
|
+
|
|
80
|
+
Four characters per token is the widely used English approximation. It
|
|
81
|
+
is only used when the upstream is silent, and it is flagged as an
|
|
82
|
+
estimate wherever it reaches a number the user sees.
|
|
83
|
+
"""
|
|
84
|
+
return max(1, len(text) // 4)
|
|
@@ -0,0 +1,238 @@
|
|
|
1
|
+
"""AWS Bedrock provider using the Converse API.
|
|
2
|
+
|
|
3
|
+
Converse is the reason this project needs no Anthropic or Meta key: one request
|
|
4
|
+
shape reaches Nova, Claude, Llama and Mistral alike. The interesting work here
|
|
5
|
+
is the translation, because Bedrock's contract differs from OpenAI's in three
|
|
6
|
+
ways that bite in production:
|
|
7
|
+
|
|
8
|
+
* system prompts are a separate top-level field, not a message;
|
|
9
|
+
* messages must strictly alternate user/assistant and must start with user;
|
|
10
|
+
* boto3 is synchronous, so every call is pushed to a worker thread and the
|
|
11
|
+
streaming iterator is bridged onto the event loop by hand.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import asyncio
|
|
17
|
+
import threading
|
|
18
|
+
from collections.abc import AsyncIterator
|
|
19
|
+
from typing import Any
|
|
20
|
+
|
|
21
|
+
import structlog
|
|
22
|
+
|
|
23
|
+
from cachellm.errors import UpstreamError
|
|
24
|
+
from cachellm.models import ChatCompletionRequest, ChatMessage
|
|
25
|
+
from cachellm.providers.base import (
|
|
26
|
+
Provider,
|
|
27
|
+
ProviderResult,
|
|
28
|
+
StreamEvent,
|
|
29
|
+
normalise_finish_reason,
|
|
30
|
+
)
|
|
31
|
+
from cachellm.settings import Settings
|
|
32
|
+
|
|
33
|
+
log = structlog.get_logger(__name__)
|
|
34
|
+
|
|
35
|
+
BEDROCK_VENDOR_PREFIXES = (
|
|
36
|
+
"amazon.",
|
|
37
|
+
"anthropic.",
|
|
38
|
+
"meta.",
|
|
39
|
+
"mistral.",
|
|
40
|
+
"cohere.",
|
|
41
|
+
"ai21.",
|
|
42
|
+
"deepseek.",
|
|
43
|
+
"qwen.",
|
|
44
|
+
"openai.",
|
|
45
|
+
"google.",
|
|
46
|
+
"moonshot",
|
|
47
|
+
"minimax.",
|
|
48
|
+
"nvidia.",
|
|
49
|
+
"zai.",
|
|
50
|
+
"writer.",
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _explain(exc: Exception) -> str:
|
|
55
|
+
"""Turn boto's less obvious failures into something actionable.
|
|
56
|
+
|
|
57
|
+
The credential one is worth special-casing: the AWS CLI's `aws login` flow
|
|
58
|
+
stores credentials through a provider that needs the optional CRT extra,
|
|
59
|
+
and boto's own message names the package without naming the install.
|
|
60
|
+
"""
|
|
61
|
+
text = str(exc)
|
|
62
|
+
if "botocore[crt]" in text or "Missing Dependency" in text:
|
|
63
|
+
return (
|
|
64
|
+
"Bedrock call failed: your AWS credentials come from `aws login`, which needs "
|
|
65
|
+
"the optional CRT extra. Install it with `uv sync --extra aws` "
|
|
66
|
+
"(or `pip install 'botocore[crt]'`), then restart the proxy. "
|
|
67
|
+
"Static keys and SSO profiles do not need this."
|
|
68
|
+
)
|
|
69
|
+
if "ExpiredToken" in text or "security token included in the request is expired" in text:
|
|
70
|
+
return "Bedrock call failed: AWS credentials have expired. Run `aws login` again."
|
|
71
|
+
if "AccessDeniedException" in text:
|
|
72
|
+
return (
|
|
73
|
+
f"Bedrock call failed: access denied. Check the model is enabled in this region "
|
|
74
|
+
f"and the identity has bedrock:InvokeModel. Original: {text[:200]}"
|
|
75
|
+
)
|
|
76
|
+
if "ThrottlingException" in text or "TooManyRequests" in text:
|
|
77
|
+
return f"Bedrock call failed: throttled by AWS. Lower concurrency and retry. {text[:160]}"
|
|
78
|
+
if "ValidationException" in text:
|
|
79
|
+
return f"Bedrock rejected the request: {text[:280]}"
|
|
80
|
+
return f"Bedrock call failed: {text[:300]}"
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def looks_like_bedrock_model(model: str) -> bool:
|
|
84
|
+
candidate = model.split("/", 1)[1] if model.startswith("bedrock/") else model
|
|
85
|
+
for region in ("us.", "eu.", "apac.", "ap."):
|
|
86
|
+
if candidate.startswith(region):
|
|
87
|
+
candidate = candidate[len(region) :]
|
|
88
|
+
break
|
|
89
|
+
return candidate.startswith(BEDROCK_VENDOR_PREFIXES)
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def to_converse_messages(messages: list[ChatMessage]) -> list[dict[str, Any]]:
|
|
93
|
+
"""Convert OpenAI messages into Converse messages, merging same-role runs."""
|
|
94
|
+
converted: list[dict[str, Any]] = []
|
|
95
|
+
for message in messages:
|
|
96
|
+
if message.role in ("system", "developer"):
|
|
97
|
+
continue
|
|
98
|
+
role = "assistant" if message.role == "assistant" else "user"
|
|
99
|
+
text = message.as_text()
|
|
100
|
+
if not text.strip():
|
|
101
|
+
continue
|
|
102
|
+
if converted and converted[-1]["role"] == role:
|
|
103
|
+
converted[-1]["content"].append({"text": text})
|
|
104
|
+
else:
|
|
105
|
+
converted.append({"role": role, "content": [{"text": text}]})
|
|
106
|
+
# Converse rejects a conversation that opens with the assistant.
|
|
107
|
+
while converted and converted[0]["role"] != "user":
|
|
108
|
+
converted.pop(0)
|
|
109
|
+
return converted
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
def build_converse_kwargs(request: ChatCompletionRequest, model_id: str) -> dict[str, Any]:
|
|
113
|
+
inference: dict[str, Any] = {}
|
|
114
|
+
if request.effective_max_tokens:
|
|
115
|
+
inference["maxTokens"] = int(request.effective_max_tokens)
|
|
116
|
+
if request.temperature is not None:
|
|
117
|
+
inference["temperature"] = float(request.temperature)
|
|
118
|
+
if request.top_p is not None:
|
|
119
|
+
inference["topP"] = float(request.top_p)
|
|
120
|
+
if request.stop:
|
|
121
|
+
stops = [request.stop] if isinstance(request.stop, str) else list(request.stop)
|
|
122
|
+
inference["stopSequences"] = stops[:4]
|
|
123
|
+
|
|
124
|
+
kwargs: dict[str, Any] = {
|
|
125
|
+
"modelId": model_id,
|
|
126
|
+
"messages": to_converse_messages(request.messages),
|
|
127
|
+
}
|
|
128
|
+
system_text = request.system_text()
|
|
129
|
+
if system_text:
|
|
130
|
+
kwargs["system"] = [{"text": system_text}]
|
|
131
|
+
if inference:
|
|
132
|
+
kwargs["inferenceConfig"] = inference
|
|
133
|
+
return kwargs
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
class BedrockProvider(Provider):
|
|
137
|
+
name = "bedrock"
|
|
138
|
+
|
|
139
|
+
def __init__(self, settings: Settings) -> None:
|
|
140
|
+
self._settings = settings
|
|
141
|
+
self._client = None
|
|
142
|
+
self._lock = threading.Lock()
|
|
143
|
+
|
|
144
|
+
def _get_client(self):
|
|
145
|
+
if self._client is None:
|
|
146
|
+
with self._lock:
|
|
147
|
+
if self._client is None:
|
|
148
|
+
import boto3
|
|
149
|
+
from botocore.config import Config
|
|
150
|
+
|
|
151
|
+
session_kwargs: dict[str, Any] = {"region_name": self._settings.aws_region}
|
|
152
|
+
if self._settings.aws_profile:
|
|
153
|
+
session_kwargs["profile_name"] = self._settings.aws_profile
|
|
154
|
+
session = boto3.Session(**session_kwargs)
|
|
155
|
+
self._client = session.client(
|
|
156
|
+
"bedrock-runtime",
|
|
157
|
+
config=Config(
|
|
158
|
+
retries={
|
|
159
|
+
"max_attempts": self._settings.provider_max_retries,
|
|
160
|
+
"mode": "standard",
|
|
161
|
+
},
|
|
162
|
+
read_timeout=int(self._settings.request_timeout),
|
|
163
|
+
connect_timeout=10,
|
|
164
|
+
),
|
|
165
|
+
)
|
|
166
|
+
return self._client
|
|
167
|
+
|
|
168
|
+
# ---------------------------------------------------------------- complete
|
|
169
|
+
async def complete(self, request: ChatCompletionRequest) -> ProviderResult:
|
|
170
|
+
model_id = self.resolve_model(request.model)
|
|
171
|
+
kwargs = build_converse_kwargs(request, model_id)
|
|
172
|
+
try:
|
|
173
|
+
response = await asyncio.to_thread(lambda: self._get_client().converse(**kwargs))
|
|
174
|
+
except Exception as exc: # boto raises many distinct client errors
|
|
175
|
+
log.warning("bedrock_error", model=model_id, error=str(exc)[:300])
|
|
176
|
+
raise UpstreamError(_explain(exc)) from exc
|
|
177
|
+
|
|
178
|
+
message = response.get("output", {}).get("message", {})
|
|
179
|
+
text = "".join(
|
|
180
|
+
block.get("text", "") for block in message.get("content", []) if "text" in block
|
|
181
|
+
)
|
|
182
|
+
usage = response.get("usage", {})
|
|
183
|
+
return ProviderResult(
|
|
184
|
+
text=text,
|
|
185
|
+
model=model_id,
|
|
186
|
+
prompt_tokens=int(usage.get("inputTokens", 0) or 0),
|
|
187
|
+
completion_tokens=int(usage.get("outputTokens", 0) or 0),
|
|
188
|
+
finish_reason=normalise_finish_reason(response.get("stopReason")),
|
|
189
|
+
raw={"metrics": response.get("metrics", {})},
|
|
190
|
+
)
|
|
191
|
+
|
|
192
|
+
# ------------------------------------------------------------------ stream
|
|
193
|
+
async def stream(self, request: ChatCompletionRequest) -> AsyncIterator[StreamEvent]:
|
|
194
|
+
model_id = self.resolve_model(request.model)
|
|
195
|
+
kwargs = build_converse_kwargs(request, model_id)
|
|
196
|
+
loop = asyncio.get_running_loop()
|
|
197
|
+
queue: asyncio.Queue[StreamEvent | Exception | None] = asyncio.Queue()
|
|
198
|
+
|
|
199
|
+
def pump() -> None:
|
|
200
|
+
"""Consume the synchronous boto EventStream on a worker thread."""
|
|
201
|
+
try:
|
|
202
|
+
response = self._get_client().converse_stream(**kwargs)
|
|
203
|
+
for event in response.get("stream", []):
|
|
204
|
+
if "contentBlockDelta" in event:
|
|
205
|
+
delta = event["contentBlockDelta"].get("delta", {}).get("text", "")
|
|
206
|
+
if delta:
|
|
207
|
+
loop.call_soon_threadsafe(queue.put_nowait, StreamEvent(delta=delta))
|
|
208
|
+
elif "messageStop" in event:
|
|
209
|
+
reason = normalise_finish_reason(event["messageStop"].get("stopReason"))
|
|
210
|
+
loop.call_soon_threadsafe(
|
|
211
|
+
queue.put_nowait, StreamEvent(finish_reason=reason)
|
|
212
|
+
)
|
|
213
|
+
elif "metadata" in event:
|
|
214
|
+
usage = event["metadata"].get("usage", {})
|
|
215
|
+
loop.call_soon_threadsafe(
|
|
216
|
+
queue.put_nowait,
|
|
217
|
+
StreamEvent(
|
|
218
|
+
prompt_tokens=int(usage.get("inputTokens", 0) or 0),
|
|
219
|
+
completion_tokens=int(usage.get("outputTokens", 0) or 0),
|
|
220
|
+
),
|
|
221
|
+
)
|
|
222
|
+
except Exception as exc: # surfaced to the consumer below
|
|
223
|
+
loop.call_soon_threadsafe(queue.put_nowait, exc)
|
|
224
|
+
finally:
|
|
225
|
+
loop.call_soon_threadsafe(queue.put_nowait, None)
|
|
226
|
+
|
|
227
|
+
task = asyncio.create_task(asyncio.to_thread(pump))
|
|
228
|
+
try:
|
|
229
|
+
while True:
|
|
230
|
+
item = await queue.get()
|
|
231
|
+
if item is None:
|
|
232
|
+
break
|
|
233
|
+
if isinstance(item, Exception):
|
|
234
|
+
log.warning("bedrock_stream_error", model=model_id, error=str(item)[:300])
|
|
235
|
+
raise UpstreamError(_explain(item)) from item
|
|
236
|
+
yield item
|
|
237
|
+
finally:
|
|
238
|
+
task.cancel()
|
|
@@ -0,0 +1,56 @@
|
|
|
1
|
+
"""Deterministic in-process provider.
|
|
2
|
+
|
|
3
|
+
Every test, the whole CI pipeline and the reproducible benchmark run against
|
|
4
|
+
this. It costs nothing, never rate-limits, and returns byte-identical answers
|
|
5
|
+
on every machine, which is what makes the published hit-rate numbers something
|
|
6
|
+
a reader can re-run rather than take on trust.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import asyncio
|
|
12
|
+
import hashlib
|
|
13
|
+
from collections.abc import AsyncIterator
|
|
14
|
+
|
|
15
|
+
from cachellm.models import ChatCompletionRequest
|
|
16
|
+
from cachellm.providers.base import Provider, ProviderResult, StreamEvent
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class FakeProvider(Provider):
|
|
20
|
+
name = "fake"
|
|
21
|
+
|
|
22
|
+
def __init__(self, latency_ms: float = 0.0, answers: dict[str, str] | None = None) -> None:
|
|
23
|
+
self.latency_ms = latency_ms
|
|
24
|
+
self.answers = answers or {}
|
|
25
|
+
self.calls = 0
|
|
26
|
+
|
|
27
|
+
def _answer(self, prompt: str) -> str:
|
|
28
|
+
if prompt in self.answers:
|
|
29
|
+
return self.answers[prompt]
|
|
30
|
+
digest = hashlib.blake2b(prompt.strip().lower().encode(), digest_size=6).hexdigest()
|
|
31
|
+
return f"[fake] answer to {prompt.strip()[:80]!r} (id {digest})"
|
|
32
|
+
|
|
33
|
+
async def complete(self, request: ChatCompletionRequest) -> ProviderResult:
|
|
34
|
+
self.calls += 1
|
|
35
|
+
if self.latency_ms:
|
|
36
|
+
await asyncio.sleep(self.latency_ms / 1000.0)
|
|
37
|
+
prompt = request.last_user_text()
|
|
38
|
+
text = self._answer(prompt)
|
|
39
|
+
return ProviderResult(
|
|
40
|
+
text=text,
|
|
41
|
+
model=self.resolve_model(request.model),
|
|
42
|
+
prompt_tokens=self.approx_tokens(prompt),
|
|
43
|
+
completion_tokens=self.approx_tokens(text),
|
|
44
|
+
finish_reason="stop",
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
async def stream(self, request: ChatCompletionRequest) -> AsyncIterator[StreamEvent]:
|
|
48
|
+
result = await self.complete(request)
|
|
49
|
+
words = result.text.split(" ")
|
|
50
|
+
for i, word in enumerate(words):
|
|
51
|
+
yield StreamEvent(delta=word if i == 0 else f" {word}")
|
|
52
|
+
yield StreamEvent(
|
|
53
|
+
finish_reason="stop",
|
|
54
|
+
prompt_tokens=result.prompt_tokens,
|
|
55
|
+
completion_tokens=result.completion_tokens,
|
|
56
|
+
)
|