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,475 @@
1
+ """Teacher-forced scoring against a self-hosted vLLM `/v1/completions` endpoint.
2
+
3
+ This is the cross-tokenizer teacher's ONLY network surface. `PromptLogprobClient.score`
4
+ submits an exact token sequence we already own and returns one logprob per position, so
5
+ the returned row indexes one for one into the teacher token ids the chunk aligner
6
+ produced. Nothing here tokenizes, renders, or samples.
7
+
8
+ Five wire facts are load-bearing (each was a real bug or a live probe finding):
9
+
10
+ 1. The prompt goes on the wire as `list[int]`, NEVER as text. vLLM re-tokenizes a text
11
+ prompt server-side with `add_special_tokens` defaulting to True, which prepends
12
+ GLM's prefix/BOS and shifts EVERY position by one against our local offsets. There
13
+ is no error: the response looks perfectly well formed and every span sum is wrong.
14
+ 2. `/v1/chat/completions` supports neither `echo` nor `prompt_logprobs`, so it cannot
15
+ score a prompt at all. The chat template is applied client-side (see
16
+ `wmo.distill.rendering`) and the rendered ids come here.
17
+ 3. Position convention: `prompt_logprobs[p]` is the distribution FOR token p (entry 0
18
+ is null because token 0 has no context). That is exactly the Tinker
19
+ `compute_logprobs` convention `wmo.distill.teacher` already uses, so `score`
20
+ returns `len(token_ids)` entries with entry 0 = None and no shifting anywhere.
21
+ A short row is rejected rather than returned: it would silently corrupt every
22
+ downstream chunk sum.
23
+ 4. The response shape varies by vLLM version: `prompt_logprobs` sits either at the top
24
+ level or under `choices[0]`. Both are read. Each position is a dict keyed by token
25
+ id (a JSON string) whose value carries a `logprob`, and the REALIZED token's entry
26
+ is the one we want, never the argmax.
27
+ 5. Auth is a Bearer token when `api_key` is set. The repo convention for self-hosted
28
+ endpoints is the `WMO_ENDPOINT_API_KEY` env var (see `wmo.providers.openai`), but
29
+ this module never reads the environment: the caller passes the key in.
30
+
31
+ Deadlines are owned here rather than through `wmo.distill.deadlines`. That module
32
+ bounds Tinker SDK calls (futures and blocking calls with no timeout parameter) and its
33
+ knobs are Tinker-named env vars; httpx already bounds an HTTP request natively, so
34
+ wrapping it in a watchdog thread would add a second, weaker timer and a misleading
35
+ `TinkerDeadlineError`. The default `timeout_s` is 1200s (20 minutes), sized for the
36
+ real workload rather than for latency headroom: merged datum length is median ~14.5k
37
+ and up to 65.5k teacher tokens, and a distillation step fires many of these
38
+ concurrently, so a request waits in the server's queue behind other prefills before
39
+ its own runs. Every request is bounded on connect, read, write, and pool, so a wedged
40
+ connection raises `PromptLogprobTimeoutError` instead of hanging a run forever.
41
+ """
42
+
43
+ from __future__ import annotations
44
+
45
+ import logging
46
+ import math
47
+ import threading
48
+ import time
49
+ from collections.abc import Callable
50
+
51
+ import httpx
52
+ from pydantic import BaseModel, ConfigDict, Field, ValidationError
53
+
54
+ from wmo.providers.base import ProviderKind, VerifyResult
55
+
56
+ logger = logging.getLogger(__name__)
57
+
58
+ DEFAULT_TIMEOUT_S = 1200.0
59
+ """Per-request wall-clock bound; see the module docstring for the sizing."""
60
+
61
+ DEFAULT_MAX_ATTEMPTS = 3
62
+ """Attempts per `score` call, counting the first; only transient failures retry."""
63
+
64
+ _RETRY_BASE_DELAY_S = 2.0
65
+ _RETRY_MAX_DELAY_S = 30.0
66
+
67
+ _RETRYABLE_STATUS = frozenset({408, 409, 425, 429, 500, 502, 503, 504})
68
+ """Statuses worth another attempt: server-side capacity, restarts, cold boots."""
69
+
70
+ _COMPLETIONS_PATH = "/v1/completions"
71
+
72
+ _VERIFY_PROBE_TOKEN_IDS: tuple[int, ...] = (1, 2, 3, 4)
73
+ """A tiny fixed sequence for verify(); small ids are valid in any real vocab."""
74
+
75
+ _MAX_CANDIDATES_IN_ERROR = 8
76
+ """How many returned candidate ids an error message quotes before eliding."""
77
+
78
+ _MAX_BODY_CHARS_IN_ERROR = 300
79
+
80
+
81
+ class PromptLogprobError(RuntimeError):
82
+ """A teacher scoring request failed, or its response was unusable.
83
+
84
+ Raised for HTTP failures, unparsable bodies, and shape violations (a row of
85
+ the wrong length, a missing realized token). The message names the endpoint
86
+ and the remedy, because a wrong shape here is silent data corruption
87
+ downstream rather than a crash.
88
+ """
89
+
90
+
91
+ class PromptLogprobTimeoutError(PromptLogprobError, TimeoutError):
92
+ """A teacher scoring request blew its wall-clock deadline.
93
+
94
+ Subclasses `TimeoutError` and keeps "timed out" in the message so retry
95
+ layers classify it as transient capacity, matching
96
+ `wmo.distill.deadlines.TinkerDeadlineError`'s contract.
97
+ """
98
+
99
+
100
+ class _LogprobEntry(BaseModel):
101
+ """One candidate in a position's `prompt_logprobs` dict (rank/decoded ignored)."""
102
+
103
+ model_config = ConfigDict(extra="ignore")
104
+
105
+ logprob: float
106
+
107
+
108
+ class _Choice(BaseModel):
109
+ """The one completion choice, which carries `prompt_logprobs` on some versions."""
110
+
111
+ model_config = ConfigDict(extra="ignore")
112
+
113
+ prompt_logprobs: list[dict[str, _LogprobEntry] | None] | None = None
114
+
115
+
116
+ class _CompletionsResponse(BaseModel):
117
+ """A `/v1/completions` response, tolerant of where `prompt_logprobs` lands."""
118
+
119
+ model_config = ConfigDict(extra="ignore")
120
+
121
+ prompt_logprobs: list[dict[str, _LogprobEntry] | None] | None = None
122
+ choices: list[_Choice] = Field(default_factory=list)
123
+
124
+
125
+ class _CompletionsRequest(BaseModel):
126
+ """The scoring request body: a token-id prompt, one throwaway sampled token."""
127
+
128
+ model_config = ConfigDict(extra="forbid")
129
+
130
+ model: str
131
+ prompt: list[int]
132
+ max_tokens: int = 1
133
+ prompt_logprobs: int = 0
134
+ temperature: float = 0.0
135
+
136
+
137
+ def _completions_url(endpoint: str) -> str:
138
+ """The `/v1/completions` URL for a configured endpoint.
139
+
140
+ Accepts either a bare server root (`https://host`) or a root that already
141
+ carries the OpenAI `/v1` prefix (the form `wmo providers` stores for
142
+ OpenAI-compatible servers), so a caller cannot accidentally produce
143
+ `/v1/v1/completions`.
144
+
145
+ Args:
146
+ endpoint: The teacher server base URL.
147
+
148
+ Returns:
149
+ The absolute URL to POST scoring requests to.
150
+
151
+ Raises:
152
+ ValueError: If `endpoint` is blank.
153
+ """
154
+ base = endpoint.strip().rstrip("/")
155
+ if not base:
156
+ raise ValueError(
157
+ "PromptLogprobClient needs a teacher endpoint URL, for example "
158
+ "'https://my-vllm-host' or 'https://my-vllm-host/v1'; got an empty string"
159
+ )
160
+ if base.endswith("/v1"):
161
+ return base + "/completions"
162
+ return base + _COMPLETIONS_PATH
163
+
164
+
165
+ class PromptLogprobClient:
166
+ """Scores exact teacher token ids on a vLLM `/v1/completions` endpoint.
167
+
168
+ One client is safe to share across threads: httpx connection pooling and the
169
+ usage counter are both synchronized, so a scoring pool can fan out over
170
+ datums against a single client (that is how the teacher is driven in a step).
171
+
172
+ Example:
173
+ >>> client = PromptLogprobClient("https://vllm-host", "zai-org/GLM-5.2")
174
+ >>> row = client.score(teacher_token_ids) # doctest: +SKIP
175
+ >>> row[0] is None # position 0 has no context # doctest: +SKIP
176
+ True
177
+ """
178
+
179
+ def __init__(
180
+ self,
181
+ endpoint: str,
182
+ model: str,
183
+ *,
184
+ api_key: str | None = None,
185
+ timeout_s: float = DEFAULT_TIMEOUT_S,
186
+ transport: httpx.BaseTransport | None = None,
187
+ max_attempts: int = DEFAULT_MAX_ATTEMPTS,
188
+ sleep: Callable[[float], None] = time.sleep,
189
+ ) -> None:
190
+ """Build a client bound to one endpoint and one served model.
191
+
192
+ Args:
193
+ endpoint: Teacher server base URL, with or without a `/v1` suffix.
194
+ model: The model id the server serves, sent as the request's `model`.
195
+ api_key: Bearer token for the endpoint. The repo convention is to pass
196
+ `WMO_ENDPOINT_API_KEY`; this module never reads the environment
197
+ itself, so the real provider keys cannot leak to an arbitrary host.
198
+ None (the norm for a private vLLM host) sends no auth header.
199
+ timeout_s: Per-request wall-clock bound in seconds, applied to connect,
200
+ read, write, and pool waits. Defaults to `DEFAULT_TIMEOUT_S`
201
+ (1200s), sized for a 65k-token prefill queued behind other
202
+ requests. Retries multiply the worst case by `max_attempts`.
203
+ transport: httpx transport override. Tests pass an
204
+ `httpx.MockTransport` so no request ever leaves the process.
205
+ max_attempts: Total attempts per `score` call. Only transient failures
206
+ (timeouts, transport errors, retryable statuses) consume attempts.
207
+ sleep: Backoff sleeper, injectable so tests do not wait.
208
+
209
+ Raises:
210
+ ValueError: If the endpoint is blank, `timeout_s` is not a positive
211
+ finite number, or `max_attempts` is below 1.
212
+ """
213
+ if not math.isfinite(timeout_s) or timeout_s <= 0:
214
+ raise ValueError(
215
+ f"timeout_s must be a positive finite number of seconds, got {timeout_s!r}; "
216
+ f"use the default ({DEFAULT_TIMEOUT_S:g}) unless the endpoint is known to be fast"
217
+ )
218
+ if max_attempts < 1:
219
+ raise ValueError(
220
+ f"max_attempts must be at least 1, got {max_attempts}; pass 1 to disable retries"
221
+ )
222
+ self._url = _completions_url(endpoint)
223
+ self._model = model
224
+ self._timeout_s = timeout_s
225
+ self._max_attempts = max_attempts
226
+ self._sleep = sleep
227
+ self._usage_tokens = 0
228
+ self._usage_lock = threading.Lock()
229
+ headers = {"Content-Type": "application/json"}
230
+ if api_key:
231
+ headers["Authorization"] = f"Bearer {api_key}"
232
+ self._client = httpx.Client(
233
+ timeout=httpx.Timeout(timeout_s),
234
+ transport=transport,
235
+ headers=headers,
236
+ )
237
+
238
+ @property
239
+ def url(self) -> str:
240
+ """The absolute scoring URL this client posts to."""
241
+ return self._url
242
+
243
+ @property
244
+ def model(self) -> str:
245
+ """The served model id sent with every request."""
246
+ return self._model
247
+
248
+ def score(self, token_ids: list[int]) -> list[float | None]:
249
+ """Teacher logprobs for an exact token sequence, one entry per position.
250
+
251
+ Args:
252
+ token_ids: The teacher's own token ids, in the teacher's vocabulary.
253
+ They are sent verbatim as integers, so the returned row aligns
254
+ index for index with this list.
255
+
256
+ Returns:
257
+ A list of `len(token_ids)` entries. Entry 0 is always None (token 0
258
+ has no context); entry p is the teacher's logprob of `token_ids[p]`
259
+ given `token_ids[:p]`.
260
+
261
+ Raises:
262
+ ValueError: If `token_ids` is empty.
263
+ PromptLogprobTimeoutError: If every attempt blew the deadline.
264
+ PromptLogprobError: On a non-retryable HTTP status, an exhausted
265
+ retry budget, an unparsable body, a row length that does not
266
+ match the prompt, or a position whose dict lacks the realized
267
+ token.
268
+ """
269
+ if not token_ids:
270
+ raise ValueError(
271
+ "score() needs at least one token id; an empty prompt has nothing to score "
272
+ "(filter empty spans out before scoring)"
273
+ )
274
+ return self._score(token_ids, count_usage=True)
275
+
276
+ def verify(self) -> VerifyResult:
277
+ """One tiny scoring probe, reporting failure as `ok=False` instead of raising.
278
+
279
+ Mirrors `wmo.providers.base.verify_via_ping` and `TinkerTeacher.verify`, so
280
+ preflight can report every misconfigured backend at once. The probe's tokens
281
+ are excluded from `usage()`.
282
+
283
+ Returns:
284
+ `ok=True` when the endpoint answered with a well-formed row, otherwise
285
+ `ok=False` with the failure text in `detail`.
286
+ """
287
+ try:
288
+ self._score(list(_VERIFY_PROBE_TOKEN_IDS), count_usage=False)
289
+ except Exception as exc: # noqa: BLE001 - verify reports failure, never raises
290
+ return VerifyResult(
291
+ ok=False, kind=ProviderKind.OPENAI, model=self._model, detail=str(exc)
292
+ )
293
+ return VerifyResult(ok=True, kind=ProviderKind.OPENAI, model=self._model, detail=self._url)
294
+
295
+ def usage(self) -> int:
296
+ """Cumulative teacher tokens submitted for scoring (verify probes excluded).
297
+
298
+ Counts every dispatched attempt, not only successful ones: a request that
299
+ timed out or died mid-response has usually already run its prefill on the
300
+ server, so the work is real. This matches `TinkerTeacher.usage`'s
301
+ "submitted, not billed-on-success" contract and feeds the same
302
+ teacher_prefill meter.
303
+ """
304
+ with self._usage_lock:
305
+ return self._usage_tokens
306
+
307
+ def close(self) -> None:
308
+ """Close the underlying connection pool."""
309
+ self._client.close()
310
+
311
+ def __enter__(self) -> PromptLogprobClient:
312
+ return self
313
+
314
+ def __exit__(self, *exc_info: object) -> None:
315
+ self.close()
316
+
317
+ def _score(self, token_ids: list[int], *, count_usage: bool) -> list[float | None]:
318
+ body = _CompletionsRequest(model=self._model, prompt=token_ids).model_dump()
319
+ response = self._post(body, token_count=len(token_ids) if count_usage else 0)
320
+ rows = self._prompt_logprob_rows(response, expected=len(token_ids))
321
+ return self._realized_row(rows, token_ids)
322
+
323
+ def _post(self, body: dict[str, object], *, token_count: int) -> _CompletionsResponse:
324
+ """POST one scoring request, retrying only transient failures."""
325
+ last_error: PromptLogprobError | None = None
326
+ for attempt in range(1, self._max_attempts + 1):
327
+ if token_count:
328
+ with self._usage_lock:
329
+ self._usage_tokens += token_count
330
+ try:
331
+ response = self._client.post(self._url, json=body)
332
+ except httpx.TimeoutException as exc:
333
+ last_error = PromptLogprobTimeoutError(
334
+ f"teacher scoring timed out after {self._timeout_s:g}s against "
335
+ f"{self._url} (attempt {attempt}/{self._max_attempts}): {exc}. Raise "
336
+ "timeout_s, lower the number of concurrent scoring calls, or check the "
337
+ "endpoint is up and not stuck in a cold boot"
338
+ )
339
+ except httpx.TransportError as exc:
340
+ last_error = PromptLogprobError(
341
+ f"teacher scoring could not reach {self._url} "
342
+ f"(attempt {attempt}/{self._max_attempts}): {exc!r}. Check the endpoint URL "
343
+ "and that the vLLM server is running and reachable from this host"
344
+ )
345
+ else:
346
+ if response.status_code < 400:
347
+ return self._parse(response)
348
+ error = self._status_error(response, attempt)
349
+ if response.status_code not in _RETRYABLE_STATUS:
350
+ raise error
351
+ last_error = error
352
+ if attempt < self._max_attempts:
353
+ delay = min(_RETRY_BASE_DELAY_S * 2 ** (attempt - 1), _RETRY_MAX_DELAY_S)
354
+ logger.warning(
355
+ "teacher scoring attempt %d/%d failed (%s); retrying in %.0fs",
356
+ attempt,
357
+ self._max_attempts,
358
+ last_error,
359
+ delay,
360
+ )
361
+ self._sleep(delay)
362
+ assert last_error is not None # noqa: S101 - the loop runs at least once
363
+ raise last_error
364
+
365
+ def _status_error(self, response: httpx.Response, attempt: int) -> PromptLogprobError:
366
+ """A typed error for a failing HTTP status, with a status-specific remedy."""
367
+ body = response.text[:_MAX_BODY_CHARS_IN_ERROR]
368
+ if response.status_code in (401, 403):
369
+ remedy = (
370
+ "the endpoint rejected the credentials: pass the api_key this server expects "
371
+ "(the repo convention is the WMO_ENDPOINT_API_KEY env var, read by the caller)"
372
+ )
373
+ elif response.status_code == 404:
374
+ remedy = (
375
+ f"no such route or model: check the base URL and that the server serves model "
376
+ f"{self._model!r} (GET /v1/models lists it)"
377
+ )
378
+ elif response.status_code in _RETRYABLE_STATUS:
379
+ remedy = (
380
+ "the server failed transiently (capacity, restart, or cold boot); retries are "
381
+ "exhausted, so lower scoring concurrency or check the server logs"
382
+ )
383
+ else:
384
+ remedy = (
385
+ "the server rejected the request: confirm this vLLM build supports "
386
+ "prompt_logprobs on /v1/completions and that no token id exceeds its vocab"
387
+ )
388
+ return PromptLogprobError(
389
+ f"teacher scoring got HTTP {response.status_code} from {self._url} "
390
+ f"(attempt {attempt}/{self._max_attempts}): {body!r}. {remedy}"
391
+ )
392
+
393
+ def _parse(self, response: httpx.Response) -> _CompletionsResponse:
394
+ """Parse a 2xx body, turning malformed JSON into a typed error."""
395
+ try:
396
+ payload = response.json()
397
+ except ValueError as exc:
398
+ raise PromptLogprobError(
399
+ f"teacher scoring got a non-JSON response from {self._url}: "
400
+ f"{response.text[:_MAX_BODY_CHARS_IN_ERROR]!r}. Check the URL points at a vLLM "
401
+ "OpenAI server and not at a proxy or web page"
402
+ ) from exc
403
+ try:
404
+ return _CompletionsResponse.model_validate(payload)
405
+ except ValidationError as exc:
406
+ raise PromptLogprobError(
407
+ f"teacher scoring could not read the response from {self._url}: {exc}. Each "
408
+ "prompt_logprobs position must be null or a dict of token id to an object with "
409
+ "a 'logprob'; check the vLLM version's response format"
410
+ ) from exc
411
+
412
+ def _prompt_logprob_rows(
413
+ self, response: _CompletionsResponse, *, expected: int
414
+ ) -> list[dict[str, _LogprobEntry] | None]:
415
+ """The per-position candidate dicts, from either shape, length-validated."""
416
+ rows = response.prompt_logprobs
417
+ if rows is None and response.choices:
418
+ rows = response.choices[0].prompt_logprobs
419
+ if rows is None:
420
+ raise PromptLogprobError(
421
+ f"teacher scoring response from {self._url} carried no prompt_logprobs (neither "
422
+ "at the top level nor under choices[0]). The request must go to /v1/completions "
423
+ "with prompt_logprobs set: /v1/chat/completions supports neither prompt_logprobs "
424
+ "nor echo, and older servers may not support it at all"
425
+ )
426
+ if len(rows) != expected:
427
+ raise PromptLogprobError(
428
+ f"teacher scoring returned {len(rows)} prompt_logprobs entries for a "
429
+ f"{expected}-token prompt at {self._url} (model {self._model}). Every position "
430
+ "must be scored, since a short or long row silently corrupts every downstream "
431
+ "span sum. Send the prompt as a list[int] (a text prompt is re-tokenized "
432
+ "server-side, which shifts every position), and confirm the server's "
433
+ f"max_model_len covers {expected} tokens"
434
+ )
435
+ return rows
436
+
437
+ def _realized_row(
438
+ self,
439
+ rows: list[dict[str, _LogprobEntry] | None],
440
+ token_ids: list[int],
441
+ ) -> list[float | None]:
442
+ """Pull each position's REALIZED token logprob, never the argmax."""
443
+ row: list[float | None] = [None]
444
+ for index in range(1, len(token_ids)):
445
+ candidates = rows[index]
446
+ token_id = token_ids[index]
447
+ entry = candidates.get(str(token_id)) if candidates else None
448
+ if entry is None:
449
+ raise PromptLogprobError(
450
+ f"teacher scoring returned no logprob for the realized token {token_id} at "
451
+ f"position {index} of {len(token_ids)} from {self._url} "
452
+ f"(candidates: {_describe_candidates(candidates)}). The realized token is "
453
+ "always included when the prompt is sent as token ids with prompt_logprobs "
454
+ f"set, so this means the ids are not in {self._model}'s vocabulary or the "
455
+ "prompt was re-tokenized server-side"
456
+ )
457
+ row.append(entry.logprob)
458
+ logger.debug(
459
+ "teacher scored %d position(s) at %s, %d tokens submitted so far",
460
+ len(token_ids),
461
+ self._url,
462
+ self.usage(),
463
+ )
464
+ return row
465
+
466
+
467
+ def _describe_candidates(candidates: dict[str, _LogprobEntry] | None) -> str:
468
+ """A short, deterministic rendering of a position's returned candidate ids."""
469
+ if not candidates:
470
+ return "none (the position was null or empty)"
471
+ keys = sorted(candidates)
472
+ shown = ", ".join(keys[:_MAX_CANDIDATES_IN_ERROR])
473
+ if len(keys) > _MAX_CANDIDATES_IN_ERROR:
474
+ return f"{len(keys)} returned, first ids {shown}, ..."
475
+ return shown