world-model-optimizer 0.2.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.
- llm_waterfall/LICENSE +21 -0
- llm_waterfall/__init__.py +53 -0
- llm_waterfall/adapters/__init__.py +36 -0
- llm_waterfall/adapters/anthropic.py +105 -0
- llm_waterfall/adapters/aws_mantle.py +47 -0
- llm_waterfall/adapters/azure_openai.py +71 -0
- llm_waterfall/adapters/base.py +51 -0
- llm_waterfall/adapters/bedrock.py +309 -0
- llm_waterfall/adapters/openai.py +130 -0
- llm_waterfall/classify.py +184 -0
- llm_waterfall/pricing.py +110 -0
- llm_waterfall/py.typed +0 -0
- llm_waterfall/types.py +295 -0
- llm_waterfall/waterfall.py +255 -0
- wmo/__init__.py +38 -0
- wmo/agents/__init__.py +7 -0
- wmo/agents/default.py +29 -0
- wmo/agents/meta.py +55 -0
- wmo/agents/optimizer.py +55 -0
- wmo/agents/project.py +928 -0
- wmo/cli/__init__.py +5 -0
- wmo/cli/agent_session.py +1123 -0
- wmo/cli/app.py +2489 -0
- wmo/cli/e2b_cmds.py +212 -0
- wmo/cli/eval_closed_loop.py +207 -0
- wmo/cli/harness_app.py +1147 -0
- wmo/cli/harness_distill.py +659 -0
- wmo/cli/hosted_session.py +880 -0
- wmo/cli/ingest_cmd.py +165 -0
- wmo/cli/model_roles.py +82 -0
- wmo/cli/platform_cmds.py +372 -0
- wmo/cli/route_app.py +274 -0
- wmo/cli/session_state.py +243 -0
- wmo/cli/ui.py +1107 -0
- wmo/cli/workspace_sync.py +504 -0
- wmo/config/__init__.py +60 -0
- wmo/config/card.py +129 -0
- wmo/config/config.py +367 -0
- wmo/config/dotenv.py +67 -0
- wmo/config/settings.py +128 -0
- wmo/config/store.py +177 -0
- wmo/conftest.py +19 -0
- wmo/connect/__init__.py +88 -0
- wmo/connect/apps.py +78 -0
- wmo/connect/brave.py +284 -0
- wmo/connect/connector.py +79 -0
- wmo/connect/credentials.py +164 -0
- wmo/connect/github.py +321 -0
- wmo/connect/google.py +627 -0
- wmo/connect/notion.py +790 -0
- wmo/connect/oauth.py +461 -0
- wmo/connect/slack.py +555 -0
- wmo/connect/store.py +199 -0
- wmo/connect/types.py +156 -0
- wmo/core/__init__.py +21 -0
- wmo/core/parsing.py +281 -0
- wmo/core/render.py +271 -0
- wmo/core/text.py +40 -0
- wmo/core/types.py +116 -0
- wmo/distill/__init__.py +14 -0
- wmo/distill/agents.py +140 -0
- wmo/distill/config.py +1006 -0
- wmo/distill/cost.py +437 -0
- wmo/distill/data.py +921 -0
- wmo/distill/deadlines.py +254 -0
- wmo/distill/fake_tinker.py +734 -0
- wmo/distill/gate.py +122 -0
- wmo/distill/loop.py +3499 -0
- wmo/distill/renderers.py +399 -0
- wmo/distill/rendering.py +620 -0
- wmo/distill/rollouts.py +726 -0
- wmo/distill/samples.py +195 -0
- wmo/distill/store.py +829 -0
- wmo/distill/teacher.py +714 -0
- wmo/distill/tokens.py +535 -0
- wmo/distill/tracking.py +552 -0
- wmo/distill/tripwire.py +411 -0
- wmo/distill/xtoken/byte_offsets.py +152 -0
- wmo/distill/xtoken/chunks.py +457 -0
- wmo/distill/xtoken/prompt_logprobs.py +475 -0
- wmo/distill/xtoken/teacher_render.py +346 -0
- wmo/engine/__init__.py +28 -0
- wmo/engine/autoconfig.py +367 -0
- wmo/engine/build.py +346 -0
- wmo/engine/demo.py +77 -0
- wmo/engine/eval_suites.py +245 -0
- wmo/engine/grounding.py +491 -0
- wmo/engine/knowledge.py +291 -0
- wmo/engine/loader.py +36 -0
- wmo/engine/play.py +92 -0
- wmo/engine/prompts.py +99 -0
- wmo/engine/replay.py +443 -0
- wmo/engine/reporting.py +58 -0
- wmo/engine/workspace.py +468 -0
- wmo/engine/world_model.py +568 -0
- wmo/env/__init__.py +22 -0
- wmo/env/base.py +121 -0
- wmo/env/closed_loop.py +229 -0
- wmo/env/episode.py +107 -0
- wmo/env/llm_agent.py +93 -0
- wmo/env/scenarios.py +73 -0
- wmo/evals/__init__.py +52 -0
- wmo/evals/agreement.py +110 -0
- wmo/evals/base.py +45 -0
- wmo/evals/closed_loop.py +480 -0
- wmo/evals/failover.py +96 -0
- wmo/evals/gold.py +127 -0
- wmo/evals/grid.py +394 -0
- wmo/evals/grid_plot.py +205 -0
- wmo/evals/harbor/__init__.py +27 -0
- wmo/evals/harbor/agent.py +573 -0
- wmo/evals/harbor/ctrf.py +171 -0
- wmo/evals/harbor/e2b_environment.py +587 -0
- wmo/evals/harbor/e2b_template_policy.py +144 -0
- wmo/evals/harbor/scorer.py +875 -0
- wmo/evals/harbor/tasks.py +140 -0
- wmo/evals/open_loop.py +194 -0
- wmo/evals/tasks.py +53 -0
- wmo/harness/__init__.py +51 -0
- wmo/harness/code_runtime.py +288 -0
- wmo/harness/create.py +1191 -0
- wmo/harness/delta.py +220 -0
- wmo/harness/doc.py +556 -0
- wmo/harness/e2b_ledger.py +342 -0
- wmo/harness/e2b_reap.py +476 -0
- wmo/harness/e2b_sandbox.py +350 -0
- wmo/harness/environment.py +35 -0
- wmo/harness/live_session.py +543 -0
- wmo/harness/mutate.py +343 -0
- wmo/harness/pi_e2b.py +1710 -0
- wmo/harness/pi_entry/entry.ts +268 -0
- wmo/harness/pi_entry/runner_frames.ts +92 -0
- wmo/harness/pi_entry/runner_live.ts +587 -0
- wmo/harness/pi_entry/runner_service.ts +270 -0
- wmo/harness/pi_entry/runner_stdio.ts +374 -0
- wmo/harness/pi_entry/runner_termination.ts +142 -0
- wmo/harness/pi_local.py +262 -0
- wmo/harness/pi_runtime.py +495 -0
- wmo/harness/pi_vendor.py +65 -0
- wmo/harness/population.py +509 -0
- wmo/harness/project_proposer.py +569 -0
- wmo/harness/proposer.py +977 -0
- wmo/harness/runner_link.py +619 -0
- wmo/harness/runtime.py +389 -0
- wmo/harness/scoring.py +247 -0
- wmo/harness/skills.py +116 -0
- wmo/harness/source_tree.py +319 -0
- wmo/harness/store.py +176 -0
- wmo/harness/tools.py +105 -0
- wmo/harness/vendor/manifest.sha256 +58 -0
- wmo/harness/vendor/pi-agent/CHANGELOG.md +556 -0
- wmo/harness/vendor/pi-agent/LICENSE +21 -0
- wmo/harness/vendor/pi-agent/README.md +488 -0
- wmo/harness/vendor/pi-agent/VENDOR.md +39 -0
- wmo/harness/vendor/pi-agent/docs/agent-harness.md +486 -0
- wmo/harness/vendor/pi-agent/docs/durable-harness.md +212 -0
- wmo/harness/vendor/pi-agent/docs/hooks.md +445 -0
- wmo/harness/vendor/pi-agent/docs/models.md +966 -0
- wmo/harness/vendor/pi-agent/docs/observability.md +376 -0
- wmo/harness/vendor/pi-agent/package.json +60 -0
- wmo/harness/vendor/pi-agent/src/agent-loop.ts +748 -0
- wmo/harness/vendor/pi-agent/src/agent.ts +575 -0
- wmo/harness/vendor/pi-agent/src/harness/agent-harness.ts +1029 -0
- wmo/harness/vendor/pi-agent/src/harness/compaction/branch-summarization.ts +261 -0
- wmo/harness/vendor/pi-agent/src/harness/compaction/compaction.ts +747 -0
- wmo/harness/vendor/pi-agent/src/harness/compaction/utils.ts +144 -0
- wmo/harness/vendor/pi-agent/src/harness/env/nodejs.ts +550 -0
- wmo/harness/vendor/pi-agent/src/harness/messages.ts +164 -0
- wmo/harness/vendor/pi-agent/src/harness/prompt-templates.ts +267 -0
- wmo/harness/vendor/pi-agent/src/harness/session/jsonl-repo.ts +177 -0
- wmo/harness/vendor/pi-agent/src/harness/session/jsonl-storage.ts +293 -0
- wmo/harness/vendor/pi-agent/src/harness/session/memory-repo.ts +50 -0
- wmo/harness/vendor/pi-agent/src/harness/session/memory-storage.ts +131 -0
- wmo/harness/vendor/pi-agent/src/harness/session/repo-utils.ts +51 -0
- wmo/harness/vendor/pi-agent/src/harness/session/session.ts +267 -0
- wmo/harness/vendor/pi-agent/src/harness/session/uuid.ts +54 -0
- wmo/harness/vendor/pi-agent/src/harness/skills.ts +375 -0
- wmo/harness/vendor/pi-agent/src/harness/system-prompt.ts +34 -0
- wmo/harness/vendor/pi-agent/src/harness/types.ts +836 -0
- wmo/harness/vendor/pi-agent/src/harness/utils/shell-output.ts +135 -0
- wmo/harness/vendor/pi-agent/src/harness/utils/truncate.ts +344 -0
- wmo/harness/vendor/pi-agent/src/index.ts +44 -0
- wmo/harness/vendor/pi-agent/src/node.ts +2 -0
- wmo/harness/vendor/pi-agent/src/proxy.ts +367 -0
- wmo/harness/vendor/pi-agent/src/types.ts +428 -0
- wmo/harness/vendor/pi-agent/test/agent-loop.test.ts +1351 -0
- wmo/harness/vendor/pi-agent/test/agent.test.ts +699 -0
- wmo/harness/vendor/pi-agent/test/e2e.test.ts +404 -0
- wmo/harness/vendor/pi-agent/test/harness/agent-harness-stream.test.ts +213 -0
- wmo/harness/vendor/pi-agent/test/harness/agent-harness.test.ts +608 -0
- wmo/harness/vendor/pi-agent/test/harness/compaction.test.ts +655 -0
- wmo/harness/vendor/pi-agent/test/harness/nodejs-env.test.ts +321 -0
- wmo/harness/vendor/pi-agent/test/harness/prompt-templates.test.ts +90 -0
- wmo/harness/vendor/pi-agent/test/harness/repo.test.ts +68 -0
- wmo/harness/vendor/pi-agent/test/harness/resource-formatting.test.ts +24 -0
- wmo/harness/vendor/pi-agent/test/harness/session-test-utils.ts +55 -0
- wmo/harness/vendor/pi-agent/test/harness/session-uuid.test.ts +50 -0
- wmo/harness/vendor/pi-agent/test/harness/session.test.ts +156 -0
- wmo/harness/vendor/pi-agent/test/harness/skills.test.ts +116 -0
- wmo/harness/vendor/pi-agent/test/harness/storage.test.ts +299 -0
- wmo/harness/vendor/pi-agent/test/harness/system-prompt.test.ts +66 -0
- wmo/harness/vendor/pi-agent/test/harness/truncate.test.ts +169 -0
- wmo/harness/vendor/pi-agent/test/scratch/simple.ts +72 -0
- wmo/harness/vendor/pi-agent/test/utils/calculate.ts +32 -0
- wmo/harness/vendor/pi-agent/test/utils/get-current-time.ts +46 -0
- wmo/harness/vendor/pi-agent/tsconfig.build.json +13 -0
- wmo/harness/vendor/pi-agent/vitest.config.ts +19 -0
- wmo/harness/vendor/pi-agent/vitest.harness.config.ts +28 -0
- wmo/harness/vendor/vendor_pi.sh +59 -0
- wmo/harness/workspace_patch.py +270 -0
- wmo/ingest/__init__.py +47 -0
- wmo/ingest/adapter.py +72 -0
- wmo/ingest/base.py +114 -0
- wmo/ingest/braintrust.py +339 -0
- wmo/ingest/detect.py +126 -0
- wmo/ingest/langfuse.py +291 -0
- wmo/ingest/langsmith.py +444 -0
- wmo/ingest/mastra.py +330 -0
- wmo/ingest/messages.py +170 -0
- wmo/ingest/normalize.py +679 -0
- wmo/ingest/otel_genai.py +69 -0
- wmo/ingest/otel_writer.py +100 -0
- wmo/ingest/phoenix.py +150 -0
- wmo/ingest/postgres.py +246 -0
- wmo/ingest/posthog.py +320 -0
- wmo/ingest/quality.py +28 -0
- wmo/ingest/stream.py +209 -0
- wmo/ingest/testdata/sample_otlp.json +60 -0
- wmo/ingest/testdata/sample_spans.jsonl +3 -0
- wmo/optimize/__init__.py +25 -0
- wmo/optimize/base.py +143 -0
- wmo/optimize/gepa.py +806 -0
- wmo/optimize/judge.py +262 -0
- wmo/optimize/judge_quality.py +359 -0
- wmo/optimize/knn.py +468 -0
- wmo/optimize/numeric.py +152 -0
- wmo/optimize/outcomes.py +103 -0
- wmo/optimize/policy.py +669 -0
- wmo/optimize/report.py +231 -0
- wmo/optimize/reward.py +129 -0
- wmo/optimize/routing.py +373 -0
- wmo/platform/__init__.py +6 -0
- wmo/platform/auth.py +115 -0
- wmo/platform/client.py +551 -0
- wmo/platform/credentials.py +126 -0
- wmo/platform/transfer.py +158 -0
- wmo/providers/__init__.py +40 -0
- wmo/providers/_bedrock_chat.py +155 -0
- wmo/providers/_openai_common.py +182 -0
- wmo/providers/_responses_common.py +472 -0
- wmo/providers/anthropic.py +134 -0
- wmo/providers/azure_openai.py +296 -0
- wmo/providers/base.py +300 -0
- wmo/providers/bedrock.py +312 -0
- wmo/providers/models.py +205 -0
- wmo/providers/openai.py +143 -0
- wmo/providers/openai_responses.py +240 -0
- wmo/providers/pool.py +170 -0
- wmo/providers/registry.py +73 -0
- wmo/providers/retry.py +151 -0
- wmo/providers/tinker.py +936 -0
- wmo/providers/waterfall.py +336 -0
- wmo/research/__init__.py +81 -0
- wmo/research/ablation.py +133 -0
- wmo/research/concurrency_plot.py +523 -0
- wmo/research/concurrency_run.py +240 -0
- wmo/research/concurrency_scaling.py +270 -0
- wmo/research/gepa_scaling.py +274 -0
- wmo/research/pipeline.py +198 -0
- wmo/research/scaling_split.py +82 -0
- wmo/research/scenario_fidelity.py +198 -0
- wmo/research/scenario_recovery.py +92 -0
- wmo/research/seed_stability.py +90 -0
- wmo/research/trace_scaling.py +348 -0
- wmo/retrieval/__init__.py +6 -0
- wmo/retrieval/embedders.py +105 -0
- wmo/retrieval/leakfree.py +52 -0
- wmo/retrieval/retriever.py +173 -0
- wmo/scenarios/__init__.py +58 -0
- wmo/scenarios/builder.py +152 -0
- wmo/scenarios/mining/__init__.py +27 -0
- wmo/scenarios/mining/clustering.py +171 -0
- wmo/scenarios/mining/facets.py +226 -0
- wmo/scenarios/mining/selection.py +220 -0
- wmo/scenarios/synthesis/__init__.py +6 -0
- wmo/scenarios/synthesis/scenario_set.py +63 -0
- wmo/scenarios/synthesis/synthesizer.py +85 -0
- wmo/scenarios/verification/__init__.py +17 -0
- wmo/scenarios/verification/judge.py +97 -0
- wmo/scenarios/verification/verify.py +135 -0
- wmo/serving/__init__.py +5 -0
- wmo/serving/builds.py +451 -0
- wmo/serving/chat.py +878 -0
- wmo/serving/endpoint_config.py +64 -0
- wmo/serving/savings.py +250 -0
- wmo/serving/server.py +553 -0
- wmo/serving/traces_source.py +206 -0
- wmo/telemetry.py +213 -0
- wmo/tracking/__init__.py +36 -0
- wmo/tracking/clock.py +24 -0
- wmo/tracking/metered.py +125 -0
- wmo/tracking/pricing.py +99 -0
- wmo/tracking/store.py +31 -0
- wmo/tracking/tracker.py +149 -0
- world_model_optimizer-0.2.0.dist-info/METADATA +203 -0
- world_model_optimizer-0.2.0.dist-info/RECORD +308 -0
- world_model_optimizer-0.2.0.dist-info/WHEEL +4 -0
- world_model_optimizer-0.2.0.dist-info/entry_points.txt +2 -0
llm_waterfall/pricing.py
ADDED
|
@@ -0,0 +1,110 @@
|
|
|
1
|
+
"""Per-model token pricing → USD cost.
|
|
2
|
+
|
|
3
|
+
Provider-agnostic: prices are keyed by a normalized model id (routing prefixes like Bedrock's
|
|
4
|
+
`us.anthropic.` are stripped before lookup), so the same Opus 4.8 row covers the direct API and
|
|
5
|
+
Bedrock. Prices are USD per 1M tokens; an unknown model costs 0.0 and `price_for` returns None so
|
|
6
|
+
callers can surface "cost unavailable" rather than silently under-reporting. Per-call overrides are
|
|
7
|
+
passed explicitly — there is no global mutable registry.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
import re
|
|
13
|
+
from collections.abc import Mapping
|
|
14
|
+
|
|
15
|
+
from pydantic import BaseModel
|
|
16
|
+
|
|
17
|
+
from llm_waterfall.types import TokenUsage
|
|
18
|
+
|
|
19
|
+
# Bedrock appends a snapshot date and/or version to the model id, e.g.
|
|
20
|
+
# `claude-haiku-4-5-20251001-v1:0` or `claude-opus-4-6-v1`. Strip them so the lookup key matches
|
|
21
|
+
# the undated table rows (`claude-haiku-4-5`). Only applied to `claude-*` ids.
|
|
22
|
+
_BEDROCK_SUFFIX = re.compile(r"(-\d{8})?(-v\d+)?(:\d+)?$")
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class ModelPrice(BaseModel):
|
|
26
|
+
"""USD per 1,000,000 tokens, split by input/output."""
|
|
27
|
+
|
|
28
|
+
input_per_mtok: float
|
|
29
|
+
output_per_mtok: float
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
# Keyed by normalized model id (see `_normalize`). USD per 1M tokens.
|
|
33
|
+
#
|
|
34
|
+
# Completion prices verified 2026-07-01 against the live vendor pricing pages (Claude via
|
|
35
|
+
# platform.claude.com models overview; OpenAI GPT-5.x Standard tier, short context). Embedding
|
|
36
|
+
# prices are long-stable list prices; treat as approximate.
|
|
37
|
+
_PRICES: dict[str, ModelPrice] = {
|
|
38
|
+
# --- Anthropic / Bedrock (Claude) ---
|
|
39
|
+
"claude-fable-5": ModelPrice(input_per_mtok=10.0, output_per_mtok=50.0),
|
|
40
|
+
"claude-mythos-5": ModelPrice(input_per_mtok=10.0, output_per_mtok=50.0),
|
|
41
|
+
"claude-opus-4-8": ModelPrice(input_per_mtok=5.0, output_per_mtok=25.0),
|
|
42
|
+
"claude-opus-4-7": ModelPrice(input_per_mtok=5.0, output_per_mtok=25.0),
|
|
43
|
+
"claude-opus-4-6": ModelPrice(input_per_mtok=5.0, output_per_mtok=25.0),
|
|
44
|
+
"claude-opus-4-5": ModelPrice(input_per_mtok=5.0, output_per_mtok=25.0),
|
|
45
|
+
"claude-opus-4-1": ModelPrice(input_per_mtok=15.0, output_per_mtok=75.0),
|
|
46
|
+
"claude-sonnet-5": ModelPrice(input_per_mtok=3.0, output_per_mtok=15.0),
|
|
47
|
+
"claude-sonnet-4-6": ModelPrice(input_per_mtok=3.0, output_per_mtok=15.0),
|
|
48
|
+
"claude-haiku-4-5": ModelPrice(input_per_mtok=1.0, output_per_mtok=5.0),
|
|
49
|
+
# --- OpenAI / Azure OpenAI (GPT-5.x; Azure deployments reuse the base model's price) ---
|
|
50
|
+
"gpt-5.5": ModelPrice(input_per_mtok=5.0, output_per_mtok=30.0),
|
|
51
|
+
"gpt-5.5-pro": ModelPrice(input_per_mtok=30.0, output_per_mtok=180.0),
|
|
52
|
+
"gpt-5.4": ModelPrice(input_per_mtok=2.5, output_per_mtok=15.0),
|
|
53
|
+
"gpt-5.4-mini": ModelPrice(input_per_mtok=0.75, output_per_mtok=4.5),
|
|
54
|
+
"gpt-5.4-nano": ModelPrice(input_per_mtok=0.2, output_per_mtok=1.25),
|
|
55
|
+
# Azure-hosted OSS deployments (qwen3-coder, agentworld, ...) are deliberately absent:
|
|
56
|
+
# a $0 placeholder row would defeat the price_for()->None "cost unavailable" contract.
|
|
57
|
+
# Supply their negotiated rates per Waterfall via the `prices` override.
|
|
58
|
+
# --- Embeddings (output tokens are always 0 for embed calls) ---
|
|
59
|
+
"text-embedding-3-small": ModelPrice(input_per_mtok=0.02, output_per_mtok=0.0),
|
|
60
|
+
"text-embedding-3-large": ModelPrice(input_per_mtok=0.13, output_per_mtok=0.0),
|
|
61
|
+
"amazon.titan-embed-text-v2:0": ModelPrice(input_per_mtok=0.02, output_per_mtok=0.0),
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def _normalize(model: str) -> str:
|
|
66
|
+
"""Strip provider/region routing prefixes so one row covers a model across providers.
|
|
67
|
+
|
|
68
|
+
Bedrock ids look like `us.anthropic.claude-opus-4-8`; the direct API uses `claude-opus-4-8`.
|
|
69
|
+
We drop a leading region segment (`us.`/`eu.`/...) and an `anthropic.` vendor segment, but keep
|
|
70
|
+
`amazon.titan-...` (its `amazon.` is part of the canonical model id, not a routing prefix).
|
|
71
|
+
"""
|
|
72
|
+
normalized = model.strip()
|
|
73
|
+
region_prefixes = ("us.", "eu.", "apac.", "us-gov.", "global.", "jp.", "au.", "ca.")
|
|
74
|
+
for prefix in region_prefixes:
|
|
75
|
+
if normalized.startswith(prefix):
|
|
76
|
+
normalized = normalized[len(prefix) :]
|
|
77
|
+
break
|
|
78
|
+
if normalized.startswith("anthropic."):
|
|
79
|
+
normalized = normalized[len("anthropic.") :]
|
|
80
|
+
if normalized.startswith("claude-"):
|
|
81
|
+
# Drop a trailing Bedrock snapshot date / version (`-20251001-v1:0`, `-v1`) so dated
|
|
82
|
+
# inference-profile ids match the undated table rows.
|
|
83
|
+
normalized = _BEDROCK_SUFFIX.sub("", normalized)
|
|
84
|
+
return normalized
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def price_for(model: str, prices: Mapping[str, ModelPrice] | None = None) -> ModelPrice | None:
|
|
88
|
+
"""The price row for `model` (after normalization), or None if unknown.
|
|
89
|
+
|
|
90
|
+
`prices` are per-caller overrides consulted before the static table; they are never merged
|
|
91
|
+
into it, so one Waterfall's overrides can't leak into another's.
|
|
92
|
+
"""
|
|
93
|
+
key = _normalize(model)
|
|
94
|
+
if prices is not None:
|
|
95
|
+
override = prices.get(key) or prices.get(model)
|
|
96
|
+
if override is not None:
|
|
97
|
+
return override
|
|
98
|
+
return _PRICES.get(key)
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def cost_usd(
|
|
102
|
+
model: str, usage: TokenUsage, prices: Mapping[str, ModelPrice] | None = None
|
|
103
|
+
) -> float:
|
|
104
|
+
"""USD cost of `usage` on `model`. Unknown models cost 0.0 (`price_for` detects that)."""
|
|
105
|
+
price = price_for(model, prices)
|
|
106
|
+
if price is None:
|
|
107
|
+
return 0.0
|
|
108
|
+
return (
|
|
109
|
+
usage.input_tokens * price.input_per_mtok + usage.output_tokens * price.output_per_mtok
|
|
110
|
+
) / 1_000_000
|
llm_waterfall/py.typed
ADDED
|
File without changes
|
llm_waterfall/types.py
ADDED
|
@@ -0,0 +1,295 @@
|
|
|
1
|
+
"""Public value types: backends, messages, per-call results, and errors.
|
|
2
|
+
|
|
3
|
+
`Backend` is a frozen dataclass (positional-friendly, hashable, safely shared across threads);
|
|
4
|
+
results are pydantic models so callers get validation and `.model_dump()` for persistence.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from collections.abc import Mapping, Sequence
|
|
10
|
+
from dataclasses import dataclass
|
|
11
|
+
from typing import Literal
|
|
12
|
+
|
|
13
|
+
from pydantic import BaseModel, ConfigDict, Field, JsonValue
|
|
14
|
+
|
|
15
|
+
Role = Literal["user", "assistant"]
|
|
16
|
+
ChatMaxTokensField = Literal["max_completion_tokens", "max_tokens"]
|
|
17
|
+
|
|
18
|
+
PROVIDERS = ("openai", "anthropic", "azure_openai", "bedrock", "aws_mantle")
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class Message(BaseModel):
|
|
22
|
+
"""One chat turn. The system prompt is a separate `complete()` param, not a message."""
|
|
23
|
+
|
|
24
|
+
role: Role
|
|
25
|
+
content: str
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
@dataclass(frozen=True)
|
|
29
|
+
class Backend:
|
|
30
|
+
"""One (provider, model, credentials) rung of the waterfall.
|
|
31
|
+
|
|
32
|
+
Credentials come from the environment (API keys) or, for Bedrock, a named AWS profile —
|
|
33
|
+
`profile` maps to `boto3.Session(profile_name=...)`, so one chain can span multiple accounts.
|
|
34
|
+
"""
|
|
35
|
+
|
|
36
|
+
provider: str # one of PROVIDERS
|
|
37
|
+
model: str
|
|
38
|
+
profile: str | None = None # bedrock: named AWS profile
|
|
39
|
+
region: str | None = None # bedrock
|
|
40
|
+
endpoint: str | None = None # azure base URL / custom OpenAI base_url
|
|
41
|
+
deployment: str | None = None # azure
|
|
42
|
+
api_version: str | None = None # azure
|
|
43
|
+
embed_model: str | None = None # None → provider default
|
|
44
|
+
embed_dim: int | None = None
|
|
45
|
+
connect_timeout_s: float = 15.0
|
|
46
|
+
# Generous read timeout: reasoning models can legitimately generate for minutes, and a
|
|
47
|
+
# mid-generation cutoff wastes the whole call — but a stalled connection must still raise
|
|
48
|
+
# (and thus fail over) instead of hanging forever.
|
|
49
|
+
read_timeout_s: float = 600.0
|
|
50
|
+
chat_max_tokens_field: ChatMaxTokensField = "max_completion_tokens"
|
|
51
|
+
|
|
52
|
+
def __post_init__(self) -> None:
|
|
53
|
+
if self.provider not in PROVIDERS:
|
|
54
|
+
raise ValueError(
|
|
55
|
+
f"unknown provider {self.provider!r}; expected one of {', '.join(PROVIDERS)}"
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
@dataclass(frozen=True)
|
|
60
|
+
class RetryPolicy:
|
|
61
|
+
"""How many times to walk the whole chain before giving up.
|
|
62
|
+
|
|
63
|
+
`rounds=1` (default) is pure failover: one attempt per backend, no sleeping. Higher values
|
|
64
|
+
wrap around — sleep with capped exponential backoff, then restart at the primary — so a long
|
|
65
|
+
unattended run survives the whole chain throttling at once. The sleep is call-local; the
|
|
66
|
+
waterfall stays stateless.
|
|
67
|
+
"""
|
|
68
|
+
|
|
69
|
+
rounds: int = 1
|
|
70
|
+
backoff_base_s: float = 15.0
|
|
71
|
+
backoff_max_s: float = 120.0
|
|
72
|
+
|
|
73
|
+
def __post_init__(self) -> None:
|
|
74
|
+
if self.rounds < 1:
|
|
75
|
+
raise ValueError("RetryPolicy.rounds must be >= 1")
|
|
76
|
+
|
|
77
|
+
def backoff_before_round(self, round_index: int) -> float:
|
|
78
|
+
"""Seconds to sleep before round `round_index` (1-based; round 1 never sleeps)."""
|
|
79
|
+
if round_index <= 1:
|
|
80
|
+
return 0.0
|
|
81
|
+
return min(self.backoff_base_s * 2 ** (round_index - 2), self.backoff_max_s)
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
class TokenUsage(BaseModel):
|
|
85
|
+
"""Raw token counts for one call (pricing converts to USD per 1M tokens)."""
|
|
86
|
+
|
|
87
|
+
input_tokens: int = 0
|
|
88
|
+
output_tokens: int = 0
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
AttemptOutcome = Literal["ok", "capacity_error", "client_error", "unsupported"]
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
class Attempt(BaseModel):
|
|
95
|
+
"""One backend try within a call — the unit of the trace the waterfall returns."""
|
|
96
|
+
|
|
97
|
+
provider: str
|
|
98
|
+
model: str
|
|
99
|
+
outcome: AttemptOutcome
|
|
100
|
+
latency_s: float
|
|
101
|
+
error: str | None = None
|
|
102
|
+
error_type: str | None = None # exception class name
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
class CompletionResult(BaseModel):
|
|
106
|
+
"""A completion plus attribution: which backend served it, what it cost, the full path."""
|
|
107
|
+
|
|
108
|
+
text: str
|
|
109
|
+
model_used: str
|
|
110
|
+
provider_used: str
|
|
111
|
+
usage: TokenUsage = Field(default_factory=TokenUsage)
|
|
112
|
+
cost_usd: float = 0.0
|
|
113
|
+
attempts: list[Attempt] = Field(default_factory=list)
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
JsonObject = dict[str, JsonValue]
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
class ChatFunctionCall(BaseModel):
|
|
120
|
+
"""Function name plus its JSON-encoded arguments in an assistant tool call."""
|
|
121
|
+
|
|
122
|
+
name: str
|
|
123
|
+
arguments: str
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
class ChatToolCall(BaseModel):
|
|
127
|
+
"""One OpenAI-compatible assistant tool call."""
|
|
128
|
+
|
|
129
|
+
id: str
|
|
130
|
+
type: Literal["function"] = "function"
|
|
131
|
+
function: ChatFunctionCall
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
class ChatMessage(BaseModel):
|
|
135
|
+
"""One structured chat turn, including tool calls and tool results."""
|
|
136
|
+
|
|
137
|
+
model_config = ConfigDict(extra="allow")
|
|
138
|
+
|
|
139
|
+
role: Literal["system", "developer", "user", "assistant", "tool"]
|
|
140
|
+
content: JsonValue = None
|
|
141
|
+
tool_calls: list[ChatToolCall] | None = None
|
|
142
|
+
tool_call_id: str | None = None
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
class ChatFunctionDefinition(BaseModel):
|
|
146
|
+
"""Function schema advertised to a tool-calling model."""
|
|
147
|
+
|
|
148
|
+
name: str
|
|
149
|
+
description: str = ""
|
|
150
|
+
parameters: JsonObject = Field(default_factory=dict)
|
|
151
|
+
strict: bool | None = None
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
class ChatTool(BaseModel):
|
|
155
|
+
"""One OpenAI-compatible function tool definition."""
|
|
156
|
+
|
|
157
|
+
type: Literal["function"] = "function"
|
|
158
|
+
function: ChatFunctionDefinition
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
class ChatRequest(BaseModel):
|
|
162
|
+
"""Provider-neutral structured chat request used by agent runtimes.
|
|
163
|
+
|
|
164
|
+
Known tool-calling fields are validated explicitly. ``extra="allow"`` preserves newer
|
|
165
|
+
OpenAI-compatible request fields emitted by an agent SDK without weakening the typed core.
|
|
166
|
+
Providers call :meth:`provider_payload` to force a non-streaming request for the framed pi
|
|
167
|
+
transport and to stamp their own routed model/deployment.
|
|
168
|
+
"""
|
|
169
|
+
|
|
170
|
+
model_config = ConfigDict(extra="allow")
|
|
171
|
+
|
|
172
|
+
messages: list[ChatMessage] = Field(default_factory=list)
|
|
173
|
+
model: str | None = None
|
|
174
|
+
tools: list[ChatTool] | None = None
|
|
175
|
+
tool_choice: JsonValue = None
|
|
176
|
+
temperature: float | None = None
|
|
177
|
+
max_tokens: int | None = None
|
|
178
|
+
max_completion_tokens: int | None = None
|
|
179
|
+
stream: bool = False
|
|
180
|
+
stream_options: JsonObject | None = None
|
|
181
|
+
|
|
182
|
+
def provider_payload(
|
|
183
|
+
self,
|
|
184
|
+
model: str,
|
|
185
|
+
*,
|
|
186
|
+
max_tokens_field: ChatMaxTokensField = "max_completion_tokens",
|
|
187
|
+
) -> JsonObject:
|
|
188
|
+
"""Return the non-streaming wire payload for a provider-routed model."""
|
|
189
|
+
payload = self.model_dump(mode="json", exclude_none=True)
|
|
190
|
+
payload["model"] = model
|
|
191
|
+
payload["stream"] = False
|
|
192
|
+
payload.pop("stream_options", None)
|
|
193
|
+
if max_tokens_field == "max_tokens":
|
|
194
|
+
alternate = payload.pop("max_completion_tokens", None)
|
|
195
|
+
if alternate is not None and "max_tokens" not in payload:
|
|
196
|
+
payload["max_tokens"] = alternate
|
|
197
|
+
else:
|
|
198
|
+
alternate = payload.pop("max_tokens", None)
|
|
199
|
+
if alternate is not None and "max_completion_tokens" not in payload:
|
|
200
|
+
payload["max_completion_tokens"] = alternate
|
|
201
|
+
return payload
|
|
202
|
+
|
|
203
|
+
|
|
204
|
+
class ChatUsage(BaseModel):
|
|
205
|
+
"""OpenAI-compatible structured completion usage."""
|
|
206
|
+
|
|
207
|
+
model_config = ConfigDict(extra="allow")
|
|
208
|
+
|
|
209
|
+
prompt_tokens: int = 0
|
|
210
|
+
completion_tokens: int = 0
|
|
211
|
+
|
|
212
|
+
|
|
213
|
+
class ChatChoice(BaseModel):
|
|
214
|
+
"""One structured completion choice."""
|
|
215
|
+
|
|
216
|
+
model_config = ConfigDict(extra="allow")
|
|
217
|
+
|
|
218
|
+
index: int = 0
|
|
219
|
+
message: ChatMessage
|
|
220
|
+
finish_reason: str | None = None
|
|
221
|
+
|
|
222
|
+
|
|
223
|
+
class ChatResponse(BaseModel):
|
|
224
|
+
"""Structured completion returned to the agent runtime."""
|
|
225
|
+
|
|
226
|
+
model_config = ConfigDict(extra="allow")
|
|
227
|
+
|
|
228
|
+
choices: list[ChatChoice]
|
|
229
|
+
usage: ChatUsage | None = None
|
|
230
|
+
model: str | None = None
|
|
231
|
+
|
|
232
|
+
def token_usage(self) -> TokenUsage:
|
|
233
|
+
"""Project provider usage onto the waterfall's canonical counters."""
|
|
234
|
+
if self.usage is None:
|
|
235
|
+
return TokenUsage()
|
|
236
|
+
return TokenUsage(
|
|
237
|
+
input_tokens=self.usage.prompt_tokens,
|
|
238
|
+
output_tokens=self.usage.completion_tokens,
|
|
239
|
+
)
|
|
240
|
+
|
|
241
|
+
def wire_payload(self) -> JsonObject:
|
|
242
|
+
"""Serialize the response back to the OpenAI-compatible pi bridge."""
|
|
243
|
+
return self.model_dump(mode="json", exclude_none=True)
|
|
244
|
+
|
|
245
|
+
|
|
246
|
+
class ChatResult(BaseModel):
|
|
247
|
+
"""A structured completion plus waterfall attribution and failover history."""
|
|
248
|
+
|
|
249
|
+
response: ChatResponse
|
|
250
|
+
model_used: str
|
|
251
|
+
provider_used: str
|
|
252
|
+
usage: TokenUsage = Field(default_factory=TokenUsage)
|
|
253
|
+
cost_usd: float = 0.0
|
|
254
|
+
attempts: list[Attempt] = Field(default_factory=list)
|
|
255
|
+
|
|
256
|
+
|
|
257
|
+
class EmbeddingResult(BaseModel):
|
|
258
|
+
"""Embedding vectors plus the same attribution as `CompletionResult`."""
|
|
259
|
+
|
|
260
|
+
vectors: list[list[float]]
|
|
261
|
+
model_used: str
|
|
262
|
+
provider_used: str
|
|
263
|
+
usage: TokenUsage = Field(default_factory=TokenUsage)
|
|
264
|
+
cost_usd: float = 0.0
|
|
265
|
+
attempts: list[Attempt] = Field(default_factory=list)
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
class VerifyResult(BaseModel):
|
|
269
|
+
"""Outcome of one backend's cheap credential/model ping."""
|
|
270
|
+
|
|
271
|
+
ok: bool
|
|
272
|
+
provider: str
|
|
273
|
+
model: str
|
|
274
|
+
detail: str = ""
|
|
275
|
+
|
|
276
|
+
|
|
277
|
+
class WaterfallExhausted(RuntimeError):
|
|
278
|
+
"""Every backend in every round was capacity-constrained. Carries the full attempt trail."""
|
|
279
|
+
|
|
280
|
+
def __init__(self, message: str, attempts: list[Attempt]) -> None:
|
|
281
|
+
super().__init__(message)
|
|
282
|
+
self.attempts = attempts
|
|
283
|
+
|
|
284
|
+
|
|
285
|
+
class EmbeddingsUnsupported(NotImplementedError):
|
|
286
|
+
"""Raised by adapters whose provider has no embeddings API; the waterfall skips them."""
|
|
287
|
+
|
|
288
|
+
|
|
289
|
+
class ToolCallingUnsupported(NotImplementedError):
|
|
290
|
+
"""Raised by adapters without structured tool-calling support; the waterfall skips them."""
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
def normalize_messages(messages: Sequence[Message | Mapping[str, str]]) -> list[Message]:
|
|
294
|
+
"""Coerce caller messages (typed or raw dicts) into the canonical `Message` list."""
|
|
295
|
+
return [m if isinstance(m, Message) else Message.model_validate(dict(m)) for m in messages]
|
|
@@ -0,0 +1,255 @@
|
|
|
1
|
+
"""The Waterfall: walk an ordered backend chain, spilling only on capacity errors.
|
|
2
|
+
|
|
3
|
+
Per call: try each backend in order. A capacity error (throttling, transient 5xx, timeout) spills
|
|
4
|
+
to the next backend; a client error (bad request, auth, validation) raises immediately — failing
|
|
5
|
+
over on those would mask a real bug behind a different model's answer. Success returns a result
|
|
6
|
+
attributed to the backend that actually served (model, provider, cost) plus the full attempt
|
|
7
|
+
trail. When every backend in every round is capacity-constrained, `WaterfallExhausted` carries
|
|
8
|
+
that trail.
|
|
9
|
+
|
|
10
|
+
Stateless by design: a `Waterfall` is immutable after construction, results are return values
|
|
11
|
+
(never side-channel logs), and the only mutable state is each adapter's lazily-built SDK client,
|
|
12
|
+
guarded by a per-adapter lock — one instance is safe to share across a thread pool.
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
from __future__ import annotations
|
|
16
|
+
|
|
17
|
+
import random
|
|
18
|
+
import time
|
|
19
|
+
from collections.abc import Callable, Mapping, Sequence
|
|
20
|
+
from typing import Literal, TypeVar
|
|
21
|
+
|
|
22
|
+
from llm_waterfall.adapters import build_adapter
|
|
23
|
+
from llm_waterfall.adapters.base import Adapter
|
|
24
|
+
from llm_waterfall.classify import outcome_for
|
|
25
|
+
from llm_waterfall.pricing import ModelPrice, cost_usd
|
|
26
|
+
from llm_waterfall.types import (
|
|
27
|
+
Attempt,
|
|
28
|
+
Backend,
|
|
29
|
+
ChatRequest,
|
|
30
|
+
ChatResult,
|
|
31
|
+
CompletionResult,
|
|
32
|
+
EmbeddingResult,
|
|
33
|
+
EmbeddingsUnsupported,
|
|
34
|
+
Message,
|
|
35
|
+
RetryPolicy,
|
|
36
|
+
TokenUsage,
|
|
37
|
+
ToolCallingUnsupported,
|
|
38
|
+
VerifyResult,
|
|
39
|
+
WaterfallExhausted,
|
|
40
|
+
normalize_messages,
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
T = TypeVar("T")
|
|
44
|
+
|
|
45
|
+
# Module-level indirection so tests can observe/skip real sleeping.
|
|
46
|
+
_sleep = time.sleep
|
|
47
|
+
|
|
48
|
+
_DEFAULT_RETRY = RetryPolicy()
|
|
49
|
+
|
|
50
|
+
_PING_MESSAGES = [Message(role="user", content="ping")]
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class Waterfall:
|
|
54
|
+
"""An immutable, thread-safe failover chain over `Backend`s."""
|
|
55
|
+
|
|
56
|
+
def __init__(
|
|
57
|
+
self,
|
|
58
|
+
backends: Sequence[Backend],
|
|
59
|
+
*,
|
|
60
|
+
retry: RetryPolicy = _DEFAULT_RETRY,
|
|
61
|
+
prices: Mapping[str, ModelPrice] | None = None,
|
|
62
|
+
adapter_factory: Callable[[Backend], Adapter] = build_adapter,
|
|
63
|
+
) -> None:
|
|
64
|
+
if not backends:
|
|
65
|
+
raise ValueError("Waterfall needs at least one backend")
|
|
66
|
+
self._backends = tuple(backends)
|
|
67
|
+
self._retry = retry
|
|
68
|
+
self._prices = dict(prices) if prices else None
|
|
69
|
+
# Adapters built eagerly (cheap — SDK clients inside are still lazy), so the tuple is
|
|
70
|
+
# immutable and there is no shared registry to guard at call time.
|
|
71
|
+
self._adapters = tuple(adapter_factory(b) for b in self._backends)
|
|
72
|
+
|
|
73
|
+
@property
|
|
74
|
+
def backends(self) -> tuple[Backend, ...]:
|
|
75
|
+
return self._backends
|
|
76
|
+
|
|
77
|
+
def complete(
|
|
78
|
+
self,
|
|
79
|
+
system: str = "",
|
|
80
|
+
messages: Sequence[Message | Mapping[str, str]] = (),
|
|
81
|
+
*,
|
|
82
|
+
temperature: float | None = None,
|
|
83
|
+
max_tokens: int = 4096,
|
|
84
|
+
) -> CompletionResult:
|
|
85
|
+
"""Run one completion down the chain.
|
|
86
|
+
|
|
87
|
+
`temperature=None` means "don't send" — current reasoning models (Claude 4.7+, GPT-5.x)
|
|
88
|
+
reject non-default sampling params, so omission is the only safe default.
|
|
89
|
+
"""
|
|
90
|
+
msgs = normalize_messages(messages)
|
|
91
|
+
|
|
92
|
+
def attempt(adapter: Adapter) -> tuple[str, TokenUsage]:
|
|
93
|
+
return adapter.complete(system, msgs, temperature=temperature, max_tokens=max_tokens)
|
|
94
|
+
|
|
95
|
+
text, usage, backend, _, attempts = self._run(attempt, unsupported=None)
|
|
96
|
+
return CompletionResult(
|
|
97
|
+
text=text,
|
|
98
|
+
model_used=backend.model,
|
|
99
|
+
provider_used=backend.provider,
|
|
100
|
+
usage=usage,
|
|
101
|
+
cost_usd=cost_usd(backend.model, usage, self._prices),
|
|
102
|
+
attempts=attempts,
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
def complete_chat(self, request: ChatRequest) -> ChatResult:
|
|
106
|
+
"""Run one structured tool-calling completion down the chain."""
|
|
107
|
+
|
|
108
|
+
def attempt(adapter: Adapter): # noqa: ANN202 - inferred from Adapter.complete_chat
|
|
109
|
+
response = adapter.complete_chat(request)
|
|
110
|
+
return response, response.token_usage()
|
|
111
|
+
|
|
112
|
+
response, usage, backend, _, attempts = self._run(
|
|
113
|
+
attempt, unsupported=ToolCallingUnsupported
|
|
114
|
+
)
|
|
115
|
+
return ChatResult(
|
|
116
|
+
response=response,
|
|
117
|
+
model_used=backend.model,
|
|
118
|
+
provider_used=backend.provider,
|
|
119
|
+
usage=usage,
|
|
120
|
+
cost_usd=cost_usd(backend.model, usage, self._prices),
|
|
121
|
+
attempts=attempts,
|
|
122
|
+
)
|
|
123
|
+
|
|
124
|
+
def embed(self, texts: Sequence[str]) -> EmbeddingResult:
|
|
125
|
+
"""Embed `texts` down the same chain; backends with no embeddings API are skipped.
|
|
126
|
+
|
|
127
|
+
Failover assumes the chain shares ONE embedding space: vectors from different embedding
|
|
128
|
+
models are not comparable, so a chain mixing embed models can silently poison a retrieval
|
|
129
|
+
index if rungs alternate mid-corpus. Keep `embed_model` consistent across rungs (e.g. the
|
|
130
|
+
same Titan model behind several AWS profiles), or embed through a single-backend chain.
|
|
131
|
+
"""
|
|
132
|
+
text_list = list(texts)
|
|
133
|
+
|
|
134
|
+
def attempt(adapter: Adapter) -> tuple[list[list[float]], TokenUsage]:
|
|
135
|
+
return adapter.embed(text_list)
|
|
136
|
+
|
|
137
|
+
vectors, usage, backend, adapter, attempts = self._run(
|
|
138
|
+
attempt, unsupported=EmbeddingsUnsupported
|
|
139
|
+
)
|
|
140
|
+
# Attribute to the model that actually embedded — the serving adapter is the single
|
|
141
|
+
# source of truth for how it resolved backend.embed_model.
|
|
142
|
+
embed_model = adapter.embed_model_id() or backend.model
|
|
143
|
+
return EmbeddingResult(
|
|
144
|
+
vectors=vectors,
|
|
145
|
+
model_used=embed_model,
|
|
146
|
+
provider_used=backend.provider,
|
|
147
|
+
usage=usage,
|
|
148
|
+
cost_usd=cost_usd(embed_model, usage, self._prices),
|
|
149
|
+
attempts=attempts,
|
|
150
|
+
)
|
|
151
|
+
|
|
152
|
+
def verify(self) -> list[VerifyResult]:
|
|
153
|
+
"""One cheap completion per backend. Reports failures, never raises.
|
|
154
|
+
|
|
155
|
+
The ping budget is 256 tokens, not 1: reasoning models (GPT-5.x) spend output tokens on
|
|
156
|
+
internal reasoning first and return 400 when the cap is hit before any visible text — a
|
|
157
|
+
1-token ping would mark a perfectly healthy backend as broken.
|
|
158
|
+
"""
|
|
159
|
+
results: list[VerifyResult] = []
|
|
160
|
+
for backend, adapter in zip(self._backends, self._adapters, strict=True):
|
|
161
|
+
try:
|
|
162
|
+
adapter.complete("", _PING_MESSAGES, temperature=None, max_tokens=256)
|
|
163
|
+
except Exception as exc: # noqa: BLE001 - verify reports failure, never raises
|
|
164
|
+
results.append(
|
|
165
|
+
VerifyResult(
|
|
166
|
+
ok=False, provider=backend.provider, model=backend.model, detail=str(exc)
|
|
167
|
+
)
|
|
168
|
+
)
|
|
169
|
+
else:
|
|
170
|
+
results.append(
|
|
171
|
+
VerifyResult(ok=True, provider=backend.provider, model=backend.model)
|
|
172
|
+
)
|
|
173
|
+
return results
|
|
174
|
+
|
|
175
|
+
def _run(
|
|
176
|
+
self,
|
|
177
|
+
attempt: Callable[[Adapter], tuple[T, TokenUsage]],
|
|
178
|
+
*,
|
|
179
|
+
unsupported: type[Exception] | None,
|
|
180
|
+
) -> tuple[T, TokenUsage, Backend, Adapter, list[Attempt]]:
|
|
181
|
+
"""The failover loop shared by complete() and embed(). All state is call-local."""
|
|
182
|
+
attempts: list[Attempt] = []
|
|
183
|
+
last_capacity_exc: Exception | None = None
|
|
184
|
+
for round_index in range(1, self._retry.rounds + 1):
|
|
185
|
+
backoff = self._retry.backoff_before_round(round_index)
|
|
186
|
+
if backoff > 0:
|
|
187
|
+
# Jittered, and never above the configured cap — callers size outer timeouts
|
|
188
|
+
# from backoff_max_s. The jitter span is reserved BELOW the cap: capping after
|
|
189
|
+
# adding jitter would collapse every concurrent caller onto exactly
|
|
190
|
+
# backoff_max_s once exponential backoff saturates, synchronizing the very
|
|
191
|
+
# retries jitter exists to spread.
|
|
192
|
+
span = 0.34 * backoff
|
|
193
|
+
base = min(backoff, self._retry.backoff_max_s - span)
|
|
194
|
+
_sleep(base + random.uniform(0, span)) # noqa: S311 - jitter
|
|
195
|
+
capacity_this_round = False
|
|
196
|
+
for backend, adapter in zip(self._backends, self._adapters, strict=True):
|
|
197
|
+
start = time.monotonic()
|
|
198
|
+
try:
|
|
199
|
+
payload, usage = attempt(adapter)
|
|
200
|
+
except (EmbeddingsUnsupported, ToolCallingUnsupported) as exc:
|
|
201
|
+
if unsupported is None or not isinstance(exc, unsupported):
|
|
202
|
+
raise
|
|
203
|
+
# Not a failure: this backend just has no embeddings API. Recorded and
|
|
204
|
+
# skipped without counting toward exhaustion.
|
|
205
|
+
attempts.append(self._attempt(backend, "unsupported", start, exc))
|
|
206
|
+
continue
|
|
207
|
+
except Exception as exc:
|
|
208
|
+
outcome = outcome_for(exc)
|
|
209
|
+
attempts.append(self._attempt(backend, outcome, start, exc))
|
|
210
|
+
if outcome == "client_error":
|
|
211
|
+
raise # a real error — never mask it behind a fallback
|
|
212
|
+
last_capacity_exc = exc
|
|
213
|
+
capacity_this_round = True
|
|
214
|
+
continue # capacity-constrained: spill to the next backend
|
|
215
|
+
attempts.append(self._attempt(backend, "ok", start, None))
|
|
216
|
+
return payload, usage, backend, adapter, attempts
|
|
217
|
+
if not capacity_this_round:
|
|
218
|
+
# Nothing transient happened this round (every backend was skipped as
|
|
219
|
+
# unsupported) — further rounds and backoff sleeps can't change the outcome.
|
|
220
|
+
break
|
|
221
|
+
if last_capacity_exc is None:
|
|
222
|
+
# Only reachable when every backend was statically unsupported; further retries
|
|
223
|
+
# cannot change that, so preserve the feature-specific configuration error.
|
|
224
|
+
if unsupported is EmbeddingsUnsupported:
|
|
225
|
+
raise EmbeddingsUnsupported(
|
|
226
|
+
"no backend in this chain supports embeddings; add a bedrock or openai "
|
|
227
|
+
"backend (anthropic has no embeddings API)."
|
|
228
|
+
)
|
|
229
|
+
if unsupported is ToolCallingUnsupported:
|
|
230
|
+
raise ToolCallingUnsupported(
|
|
231
|
+
"no backend in this chain supports structured tool calling; add an openai, "
|
|
232
|
+
"azure_openai, or bedrock backend."
|
|
233
|
+
)
|
|
234
|
+
raise AssertionError("waterfall ended without a result or capacity error")
|
|
235
|
+
message = (
|
|
236
|
+
f"every backend was capacity-constrained after {len(attempts)} attempts "
|
|
237
|
+
f"across {self._retry.rounds} round(s)"
|
|
238
|
+
)
|
|
239
|
+
raise WaterfallExhausted(message, attempts) from last_capacity_exc
|
|
240
|
+
|
|
241
|
+
@staticmethod
|
|
242
|
+
def _attempt(
|
|
243
|
+
backend: Backend,
|
|
244
|
+
outcome: Literal["ok", "capacity_error", "client_error", "unsupported"],
|
|
245
|
+
start: float,
|
|
246
|
+
exc: Exception | None,
|
|
247
|
+
) -> Attempt:
|
|
248
|
+
return Attempt(
|
|
249
|
+
provider=backend.provider,
|
|
250
|
+
model=backend.model,
|
|
251
|
+
outcome=outcome,
|
|
252
|
+
latency_s=time.monotonic() - start,
|
|
253
|
+
error=str(exc) if exc is not None else None,
|
|
254
|
+
error_type=type(exc).__name__ if exc is not None else None,
|
|
255
|
+
)
|