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.
Files changed (308) hide show
  1. llm_waterfall/LICENSE +21 -0
  2. llm_waterfall/__init__.py +53 -0
  3. llm_waterfall/adapters/__init__.py +36 -0
  4. llm_waterfall/adapters/anthropic.py +105 -0
  5. llm_waterfall/adapters/aws_mantle.py +47 -0
  6. llm_waterfall/adapters/azure_openai.py +71 -0
  7. llm_waterfall/adapters/base.py +51 -0
  8. llm_waterfall/adapters/bedrock.py +309 -0
  9. llm_waterfall/adapters/openai.py +130 -0
  10. llm_waterfall/classify.py +184 -0
  11. llm_waterfall/pricing.py +110 -0
  12. llm_waterfall/py.typed +0 -0
  13. llm_waterfall/types.py +295 -0
  14. llm_waterfall/waterfall.py +255 -0
  15. wmo/__init__.py +38 -0
  16. wmo/agents/__init__.py +7 -0
  17. wmo/agents/default.py +29 -0
  18. wmo/agents/meta.py +55 -0
  19. wmo/agents/optimizer.py +55 -0
  20. wmo/agents/project.py +928 -0
  21. wmo/cli/__init__.py +5 -0
  22. wmo/cli/agent_session.py +1123 -0
  23. wmo/cli/app.py +2489 -0
  24. wmo/cli/e2b_cmds.py +212 -0
  25. wmo/cli/eval_closed_loop.py +207 -0
  26. wmo/cli/harness_app.py +1147 -0
  27. wmo/cli/harness_distill.py +659 -0
  28. wmo/cli/hosted_session.py +880 -0
  29. wmo/cli/ingest_cmd.py +165 -0
  30. wmo/cli/model_roles.py +82 -0
  31. wmo/cli/platform_cmds.py +372 -0
  32. wmo/cli/route_app.py +274 -0
  33. wmo/cli/session_state.py +243 -0
  34. wmo/cli/ui.py +1107 -0
  35. wmo/cli/workspace_sync.py +504 -0
  36. wmo/config/__init__.py +60 -0
  37. wmo/config/card.py +129 -0
  38. wmo/config/config.py +367 -0
  39. wmo/config/dotenv.py +67 -0
  40. wmo/config/settings.py +128 -0
  41. wmo/config/store.py +177 -0
  42. wmo/conftest.py +19 -0
  43. wmo/connect/__init__.py +88 -0
  44. wmo/connect/apps.py +78 -0
  45. wmo/connect/brave.py +284 -0
  46. wmo/connect/connector.py +79 -0
  47. wmo/connect/credentials.py +164 -0
  48. wmo/connect/github.py +321 -0
  49. wmo/connect/google.py +627 -0
  50. wmo/connect/notion.py +790 -0
  51. wmo/connect/oauth.py +461 -0
  52. wmo/connect/slack.py +555 -0
  53. wmo/connect/store.py +199 -0
  54. wmo/connect/types.py +156 -0
  55. wmo/core/__init__.py +21 -0
  56. wmo/core/parsing.py +281 -0
  57. wmo/core/render.py +271 -0
  58. wmo/core/text.py +40 -0
  59. wmo/core/types.py +116 -0
  60. wmo/distill/__init__.py +14 -0
  61. wmo/distill/agents.py +140 -0
  62. wmo/distill/config.py +1006 -0
  63. wmo/distill/cost.py +437 -0
  64. wmo/distill/data.py +921 -0
  65. wmo/distill/deadlines.py +254 -0
  66. wmo/distill/fake_tinker.py +734 -0
  67. wmo/distill/gate.py +122 -0
  68. wmo/distill/loop.py +3499 -0
  69. wmo/distill/renderers.py +399 -0
  70. wmo/distill/rendering.py +620 -0
  71. wmo/distill/rollouts.py +726 -0
  72. wmo/distill/samples.py +195 -0
  73. wmo/distill/store.py +829 -0
  74. wmo/distill/teacher.py +714 -0
  75. wmo/distill/tokens.py +535 -0
  76. wmo/distill/tracking.py +552 -0
  77. wmo/distill/tripwire.py +411 -0
  78. wmo/distill/xtoken/byte_offsets.py +152 -0
  79. wmo/distill/xtoken/chunks.py +457 -0
  80. wmo/distill/xtoken/prompt_logprobs.py +475 -0
  81. wmo/distill/xtoken/teacher_render.py +346 -0
  82. wmo/engine/__init__.py +28 -0
  83. wmo/engine/autoconfig.py +367 -0
  84. wmo/engine/build.py +346 -0
  85. wmo/engine/demo.py +77 -0
  86. wmo/engine/eval_suites.py +245 -0
  87. wmo/engine/grounding.py +491 -0
  88. wmo/engine/knowledge.py +291 -0
  89. wmo/engine/loader.py +36 -0
  90. wmo/engine/play.py +92 -0
  91. wmo/engine/prompts.py +99 -0
  92. wmo/engine/replay.py +443 -0
  93. wmo/engine/reporting.py +58 -0
  94. wmo/engine/workspace.py +468 -0
  95. wmo/engine/world_model.py +568 -0
  96. wmo/env/__init__.py +22 -0
  97. wmo/env/base.py +121 -0
  98. wmo/env/closed_loop.py +229 -0
  99. wmo/env/episode.py +107 -0
  100. wmo/env/llm_agent.py +93 -0
  101. wmo/env/scenarios.py +73 -0
  102. wmo/evals/__init__.py +52 -0
  103. wmo/evals/agreement.py +110 -0
  104. wmo/evals/base.py +45 -0
  105. wmo/evals/closed_loop.py +480 -0
  106. wmo/evals/failover.py +96 -0
  107. wmo/evals/gold.py +127 -0
  108. wmo/evals/grid.py +394 -0
  109. wmo/evals/grid_plot.py +205 -0
  110. wmo/evals/harbor/__init__.py +27 -0
  111. wmo/evals/harbor/agent.py +573 -0
  112. wmo/evals/harbor/ctrf.py +171 -0
  113. wmo/evals/harbor/e2b_environment.py +587 -0
  114. wmo/evals/harbor/e2b_template_policy.py +144 -0
  115. wmo/evals/harbor/scorer.py +875 -0
  116. wmo/evals/harbor/tasks.py +140 -0
  117. wmo/evals/open_loop.py +194 -0
  118. wmo/evals/tasks.py +53 -0
  119. wmo/harness/__init__.py +51 -0
  120. wmo/harness/code_runtime.py +288 -0
  121. wmo/harness/create.py +1191 -0
  122. wmo/harness/delta.py +220 -0
  123. wmo/harness/doc.py +556 -0
  124. wmo/harness/e2b_ledger.py +342 -0
  125. wmo/harness/e2b_reap.py +476 -0
  126. wmo/harness/e2b_sandbox.py +350 -0
  127. wmo/harness/environment.py +35 -0
  128. wmo/harness/live_session.py +543 -0
  129. wmo/harness/mutate.py +343 -0
  130. wmo/harness/pi_e2b.py +1710 -0
  131. wmo/harness/pi_entry/entry.ts +268 -0
  132. wmo/harness/pi_entry/runner_frames.ts +92 -0
  133. wmo/harness/pi_entry/runner_live.ts +587 -0
  134. wmo/harness/pi_entry/runner_service.ts +270 -0
  135. wmo/harness/pi_entry/runner_stdio.ts +374 -0
  136. wmo/harness/pi_entry/runner_termination.ts +142 -0
  137. wmo/harness/pi_local.py +262 -0
  138. wmo/harness/pi_runtime.py +495 -0
  139. wmo/harness/pi_vendor.py +65 -0
  140. wmo/harness/population.py +509 -0
  141. wmo/harness/project_proposer.py +569 -0
  142. wmo/harness/proposer.py +977 -0
  143. wmo/harness/runner_link.py +619 -0
  144. wmo/harness/runtime.py +389 -0
  145. wmo/harness/scoring.py +247 -0
  146. wmo/harness/skills.py +116 -0
  147. wmo/harness/source_tree.py +319 -0
  148. wmo/harness/store.py +176 -0
  149. wmo/harness/tools.py +105 -0
  150. wmo/harness/vendor/manifest.sha256 +58 -0
  151. wmo/harness/vendor/pi-agent/CHANGELOG.md +556 -0
  152. wmo/harness/vendor/pi-agent/LICENSE +21 -0
  153. wmo/harness/vendor/pi-agent/README.md +488 -0
  154. wmo/harness/vendor/pi-agent/VENDOR.md +39 -0
  155. wmo/harness/vendor/pi-agent/docs/agent-harness.md +486 -0
  156. wmo/harness/vendor/pi-agent/docs/durable-harness.md +212 -0
  157. wmo/harness/vendor/pi-agent/docs/hooks.md +445 -0
  158. wmo/harness/vendor/pi-agent/docs/models.md +966 -0
  159. wmo/harness/vendor/pi-agent/docs/observability.md +376 -0
  160. wmo/harness/vendor/pi-agent/package.json +60 -0
  161. wmo/harness/vendor/pi-agent/src/agent-loop.ts +748 -0
  162. wmo/harness/vendor/pi-agent/src/agent.ts +575 -0
  163. wmo/harness/vendor/pi-agent/src/harness/agent-harness.ts +1029 -0
  164. wmo/harness/vendor/pi-agent/src/harness/compaction/branch-summarization.ts +261 -0
  165. wmo/harness/vendor/pi-agent/src/harness/compaction/compaction.ts +747 -0
  166. wmo/harness/vendor/pi-agent/src/harness/compaction/utils.ts +144 -0
  167. wmo/harness/vendor/pi-agent/src/harness/env/nodejs.ts +550 -0
  168. wmo/harness/vendor/pi-agent/src/harness/messages.ts +164 -0
  169. wmo/harness/vendor/pi-agent/src/harness/prompt-templates.ts +267 -0
  170. wmo/harness/vendor/pi-agent/src/harness/session/jsonl-repo.ts +177 -0
  171. wmo/harness/vendor/pi-agent/src/harness/session/jsonl-storage.ts +293 -0
  172. wmo/harness/vendor/pi-agent/src/harness/session/memory-repo.ts +50 -0
  173. wmo/harness/vendor/pi-agent/src/harness/session/memory-storage.ts +131 -0
  174. wmo/harness/vendor/pi-agent/src/harness/session/repo-utils.ts +51 -0
  175. wmo/harness/vendor/pi-agent/src/harness/session/session.ts +267 -0
  176. wmo/harness/vendor/pi-agent/src/harness/session/uuid.ts +54 -0
  177. wmo/harness/vendor/pi-agent/src/harness/skills.ts +375 -0
  178. wmo/harness/vendor/pi-agent/src/harness/system-prompt.ts +34 -0
  179. wmo/harness/vendor/pi-agent/src/harness/types.ts +836 -0
  180. wmo/harness/vendor/pi-agent/src/harness/utils/shell-output.ts +135 -0
  181. wmo/harness/vendor/pi-agent/src/harness/utils/truncate.ts +344 -0
  182. wmo/harness/vendor/pi-agent/src/index.ts +44 -0
  183. wmo/harness/vendor/pi-agent/src/node.ts +2 -0
  184. wmo/harness/vendor/pi-agent/src/proxy.ts +367 -0
  185. wmo/harness/vendor/pi-agent/src/types.ts +428 -0
  186. wmo/harness/vendor/pi-agent/test/agent-loop.test.ts +1351 -0
  187. wmo/harness/vendor/pi-agent/test/agent.test.ts +699 -0
  188. wmo/harness/vendor/pi-agent/test/e2e.test.ts +404 -0
  189. wmo/harness/vendor/pi-agent/test/harness/agent-harness-stream.test.ts +213 -0
  190. wmo/harness/vendor/pi-agent/test/harness/agent-harness.test.ts +608 -0
  191. wmo/harness/vendor/pi-agent/test/harness/compaction.test.ts +655 -0
  192. wmo/harness/vendor/pi-agent/test/harness/nodejs-env.test.ts +321 -0
  193. wmo/harness/vendor/pi-agent/test/harness/prompt-templates.test.ts +90 -0
  194. wmo/harness/vendor/pi-agent/test/harness/repo.test.ts +68 -0
  195. wmo/harness/vendor/pi-agent/test/harness/resource-formatting.test.ts +24 -0
  196. wmo/harness/vendor/pi-agent/test/harness/session-test-utils.ts +55 -0
  197. wmo/harness/vendor/pi-agent/test/harness/session-uuid.test.ts +50 -0
  198. wmo/harness/vendor/pi-agent/test/harness/session.test.ts +156 -0
  199. wmo/harness/vendor/pi-agent/test/harness/skills.test.ts +116 -0
  200. wmo/harness/vendor/pi-agent/test/harness/storage.test.ts +299 -0
  201. wmo/harness/vendor/pi-agent/test/harness/system-prompt.test.ts +66 -0
  202. wmo/harness/vendor/pi-agent/test/harness/truncate.test.ts +169 -0
  203. wmo/harness/vendor/pi-agent/test/scratch/simple.ts +72 -0
  204. wmo/harness/vendor/pi-agent/test/utils/calculate.ts +32 -0
  205. wmo/harness/vendor/pi-agent/test/utils/get-current-time.ts +46 -0
  206. wmo/harness/vendor/pi-agent/tsconfig.build.json +13 -0
  207. wmo/harness/vendor/pi-agent/vitest.config.ts +19 -0
  208. wmo/harness/vendor/pi-agent/vitest.harness.config.ts +28 -0
  209. wmo/harness/vendor/vendor_pi.sh +59 -0
  210. wmo/harness/workspace_patch.py +270 -0
  211. wmo/ingest/__init__.py +47 -0
  212. wmo/ingest/adapter.py +72 -0
  213. wmo/ingest/base.py +114 -0
  214. wmo/ingest/braintrust.py +339 -0
  215. wmo/ingest/detect.py +126 -0
  216. wmo/ingest/langfuse.py +291 -0
  217. wmo/ingest/langsmith.py +444 -0
  218. wmo/ingest/mastra.py +330 -0
  219. wmo/ingest/messages.py +170 -0
  220. wmo/ingest/normalize.py +679 -0
  221. wmo/ingest/otel_genai.py +69 -0
  222. wmo/ingest/otel_writer.py +100 -0
  223. wmo/ingest/phoenix.py +150 -0
  224. wmo/ingest/postgres.py +246 -0
  225. wmo/ingest/posthog.py +320 -0
  226. wmo/ingest/quality.py +28 -0
  227. wmo/ingest/stream.py +209 -0
  228. wmo/ingest/testdata/sample_otlp.json +60 -0
  229. wmo/ingest/testdata/sample_spans.jsonl +3 -0
  230. wmo/optimize/__init__.py +25 -0
  231. wmo/optimize/base.py +143 -0
  232. wmo/optimize/gepa.py +806 -0
  233. wmo/optimize/judge.py +262 -0
  234. wmo/optimize/judge_quality.py +359 -0
  235. wmo/optimize/knn.py +468 -0
  236. wmo/optimize/numeric.py +152 -0
  237. wmo/optimize/outcomes.py +103 -0
  238. wmo/optimize/policy.py +669 -0
  239. wmo/optimize/report.py +231 -0
  240. wmo/optimize/reward.py +129 -0
  241. wmo/optimize/routing.py +373 -0
  242. wmo/platform/__init__.py +6 -0
  243. wmo/platform/auth.py +115 -0
  244. wmo/platform/client.py +551 -0
  245. wmo/platform/credentials.py +126 -0
  246. wmo/platform/transfer.py +158 -0
  247. wmo/providers/__init__.py +40 -0
  248. wmo/providers/_bedrock_chat.py +155 -0
  249. wmo/providers/_openai_common.py +182 -0
  250. wmo/providers/_responses_common.py +472 -0
  251. wmo/providers/anthropic.py +134 -0
  252. wmo/providers/azure_openai.py +296 -0
  253. wmo/providers/base.py +300 -0
  254. wmo/providers/bedrock.py +312 -0
  255. wmo/providers/models.py +205 -0
  256. wmo/providers/openai.py +143 -0
  257. wmo/providers/openai_responses.py +240 -0
  258. wmo/providers/pool.py +170 -0
  259. wmo/providers/registry.py +73 -0
  260. wmo/providers/retry.py +151 -0
  261. wmo/providers/tinker.py +936 -0
  262. wmo/providers/waterfall.py +336 -0
  263. wmo/research/__init__.py +81 -0
  264. wmo/research/ablation.py +133 -0
  265. wmo/research/concurrency_plot.py +523 -0
  266. wmo/research/concurrency_run.py +240 -0
  267. wmo/research/concurrency_scaling.py +270 -0
  268. wmo/research/gepa_scaling.py +274 -0
  269. wmo/research/pipeline.py +198 -0
  270. wmo/research/scaling_split.py +82 -0
  271. wmo/research/scenario_fidelity.py +198 -0
  272. wmo/research/scenario_recovery.py +92 -0
  273. wmo/research/seed_stability.py +90 -0
  274. wmo/research/trace_scaling.py +348 -0
  275. wmo/retrieval/__init__.py +6 -0
  276. wmo/retrieval/embedders.py +105 -0
  277. wmo/retrieval/leakfree.py +52 -0
  278. wmo/retrieval/retriever.py +173 -0
  279. wmo/scenarios/__init__.py +58 -0
  280. wmo/scenarios/builder.py +152 -0
  281. wmo/scenarios/mining/__init__.py +27 -0
  282. wmo/scenarios/mining/clustering.py +171 -0
  283. wmo/scenarios/mining/facets.py +226 -0
  284. wmo/scenarios/mining/selection.py +220 -0
  285. wmo/scenarios/synthesis/__init__.py +6 -0
  286. wmo/scenarios/synthesis/scenario_set.py +63 -0
  287. wmo/scenarios/synthesis/synthesizer.py +85 -0
  288. wmo/scenarios/verification/__init__.py +17 -0
  289. wmo/scenarios/verification/judge.py +97 -0
  290. wmo/scenarios/verification/verify.py +135 -0
  291. wmo/serving/__init__.py +5 -0
  292. wmo/serving/builds.py +451 -0
  293. wmo/serving/chat.py +878 -0
  294. wmo/serving/endpoint_config.py +64 -0
  295. wmo/serving/savings.py +250 -0
  296. wmo/serving/server.py +553 -0
  297. wmo/serving/traces_source.py +206 -0
  298. wmo/telemetry.py +213 -0
  299. wmo/tracking/__init__.py +36 -0
  300. wmo/tracking/clock.py +24 -0
  301. wmo/tracking/metered.py +125 -0
  302. wmo/tracking/pricing.py +99 -0
  303. wmo/tracking/store.py +31 -0
  304. wmo/tracking/tracker.py +149 -0
  305. world_model_optimizer-0.2.0.dist-info/METADATA +203 -0
  306. world_model_optimizer-0.2.0.dist-info/RECORD +308 -0
  307. world_model_optimizer-0.2.0.dist-info/WHEEL +4 -0
  308. world_model_optimizer-0.2.0.dist-info/entry_points.txt +2 -0
@@ -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
+ )