dev-double 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- dev_double/__init__.py +3 -0
- dev_double/apps/__init__.py +1 -0
- dev_double/apps/cli/__init__.py +5 -0
- dev_double/apps/cli/main.py +122 -0
- dev_double/apps/client/__init__.py +5 -0
- dev_double/apps/client/http_client.py +102 -0
- dev_double/apps/composition.py +71 -0
- dev_double/apps/config.py +36 -0
- dev_double/apps/server/__init__.py +5 -0
- dev_double/apps/server/app.py +122 -0
- dev_double/apps/server/recorder.py +62 -0
- dev_double/apps/server/transport/__init__.py +39 -0
- dev_double/apps/server/transport/choice_use_cases.py +121 -0
- dev_double/apps/server/transport/common.py +37 -0
- dev_double/apps/server/transport/decide.py +131 -0
- dev_double/apps/server/transport/extract.py +93 -0
- dev_double/apps/server/transport/guard_judge.py +97 -0
- dev_double/apps/server/transport/rerank.py +55 -0
- dev_double/client.py +5 -0
- dev_double/core/__init__.py +1 -0
- dev_double/core/decision/__init__.py +62 -0
- dev_double/core/decision/answer_shaping.py +48 -0
- dev_double/core/decision/confidence.py +19 -0
- dev_double/core/decision/decider_basic_impl.py +65 -0
- dev_double/core/decision/decision_service_basic_impl.py +141 -0
- dev_double/core/decision/defaults.py +40 -0
- dev_double/core/decision/errors.py +7 -0
- dev_double/core/decision/extract_questions.py +40 -0
- dev_double/core/decision/field_extraction.py +53 -0
- dev_double/core/decision/i_clock.py +11 -0
- dev_double/core/decision/i_decider.py +36 -0
- dev_double/core/decision/i_decision_service.py +40 -0
- dev_double/core/decision/i_engine.py +27 -0
- dev_double/core/decision/i_extractor.py +19 -0
- dev_double/core/decision/i_generator.py +19 -0
- dev_double/core/decision/i_id_provider.py +9 -0
- dev_double/core/decision/i_record_reader.py +25 -0
- dev_double/core/decision/i_reranker.py +18 -0
- dev_double/core/decision/label_scoring.py +89 -0
- dev_double/core/decision/prompts.py +130 -0
- dev_double/core/decision/record_parsing.py +153 -0
- dev_double/core/decision/t_answer.py +32 -0
- dev_double/core/decision/t_classify.py +42 -0
- dev_double/core/decision/t_decide.py +26 -0
- dev_double/core/decision/t_extract.py +89 -0
- dev_double/core/decision/t_gate.py +36 -0
- dev_double/core/decision/t_generate.py +37 -0
- dev_double/core/decision/t_guard.py +38 -0
- dev_double/core/decision/t_input.py +8 -0
- dev_double/core/decision/t_judge.py +35 -0
- dev_double/core/decision/t_label_query.py +36 -0
- dev_double/core/decision/t_meta.py +17 -0
- dev_double/core/decision/t_question.py +70 -0
- dev_double/core/decision/t_rerank.py +36 -0
- dev_double/core/decision/t_route.py +28 -0
- dev_double/core/decision/t_usage.py +15 -0
- dev_double/core/decision/tracker.py +27 -0
- dev_double/core/decision/use_case_questions.py +80 -0
- dev_double/core/decision/value_parsing.py +109 -0
- dev_double/providers/__init__.py +1 -0
- dev_double/providers/mock/__init__.py +1 -0
- dev_double/providers/mock/decision/__init__.py +5 -0
- dev_double/providers/mock/decision/clock_mock_impl.py +19 -0
- dev_double/providers/mock/decision/engine_mock_impl.py +78 -0
- dev_double/providers/mock/decision/id_provider_mock_impl.py +15 -0
- dev_double/providers/needle/__init__.py +1 -0
- dev_double/providers/needle/decision/__init__.py +4 -0
- dev_double/providers/needle/decision/decider_needle_impl.py +117 -0
- dev_double/providers/needle/decision/record_tool.py +75 -0
- dev_double/providers/openai/__init__.py +1 -0
- dev_double/providers/openai/decision/__init__.py +3 -0
- dev_double/providers/openai/decision/engine_openai_impl.py +154 -0
- dev_double/providers/std/__init__.py +1 -0
- dev_double/providers/std/decision/__init__.py +4 -0
- dev_double/providers/std/decision/clock_std_impl.py +12 -0
- dev_double/providers/std/decision/id_provider_std_impl.py +12 -0
- dev_double/providers/systemone/__init__.py +1 -0
- dev_double/providers/systemone/decision/__init__.py +3 -0
- dev_double/providers/systemone/decision/decider_system_one_impl.py +100 -0
- dev_double-0.1.0.dist-info/METADATA +354 -0
- dev_double-0.1.0.dist-info/RECORD +84 -0
- dev_double-0.1.0.dist-info/WHEEL +4 -0
- dev_double-0.1.0.dist-info/entry_points.txt +2 -0
- dev_double-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
"""Extraction as one Needle tool call: every field becomes an argument
|
|
2
|
+
of a ``record`` tool, and the call's arguments are mapped back to typed values.
|
|
3
|
+
|
|
4
|
+
pydantic is imported lazily, like needle, so this module imports without either.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from typing import Any, Literal, Optional
|
|
10
|
+
|
|
11
|
+
from dev_double.core.decision.confidence import DIGITS
|
|
12
|
+
from dev_double.core.decision.t_extract import EMPTY, TExtractField, TFieldValue
|
|
13
|
+
from dev_double.core.decision.tracker import Tracker
|
|
14
|
+
from dev_double.core.decision.value_parsing import coerce_json_value
|
|
15
|
+
|
|
16
|
+
_PY_TYPES: dict[str, type] = {"string": str, "number": float, "integer": int, "boolean": bool}
|
|
17
|
+
|
|
18
|
+
RECORD_SYSTEM = (
|
|
19
|
+
"Extract the described fields from the input. Always call record once. "
|
|
20
|
+
"Leave out a text or number field the input does not contain."
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def fields_key(fields: dict[str, TExtractField]) -> tuple:
|
|
25
|
+
"""A hashable key for a field set, to reuse one agent per distinct set."""
|
|
26
|
+
return tuple((n, f.type, f.description, tuple((f.options or {}).items())) for n, f in fields.items()) # type: ignore[union-attr]
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def record_model(fields: dict[str, TExtractField]) -> Any:
|
|
30
|
+
"""The tool schema: field names are aliases, so any name (even "due date") works."""
|
|
31
|
+
from pydantic import Field, create_model
|
|
32
|
+
|
|
33
|
+
specs: dict[str, Any] = {}
|
|
34
|
+
for i, (name, f) in enumerate(fields.items()):
|
|
35
|
+
description = f.description or name
|
|
36
|
+
if f.type == "enum":
|
|
37
|
+
assert isinstance(f.options, dict)
|
|
38
|
+
kind: Any = Literal[tuple(f.options)] # type: ignore[valid-type]
|
|
39
|
+
description += " One of: " + "; ".join(f"{k}: {v}" for k, v in f.options.items())
|
|
40
|
+
else:
|
|
41
|
+
kind = _PY_TYPES[f.type]
|
|
42
|
+
if f.type in ("enum", "boolean"):
|
|
43
|
+
# Always answerable from the text, so required: optional fields with no
|
|
44
|
+
# matching span come back empty from Needle.
|
|
45
|
+
specs[f"f{i}"] = (kind, Field(..., alias=name, description=description))
|
|
46
|
+
else:
|
|
47
|
+
specs[f"f{i}"] = (Optional[kind], Field(None, alias=name, description=description))
|
|
48
|
+
return create_model("record", __doc__="Record the fields found in the input.", **specs)
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _coerce(f: TExtractField, raw: Any) -> tuple[Any, Optional[str]]:
|
|
52
|
+
if f.type == "boolean":
|
|
53
|
+
return (raw, None) if isinstance(raw, bool) else (None, f"not a boolean: {raw!r}")
|
|
54
|
+
if f.type == "enum":
|
|
55
|
+
return (raw, None) if raw in (f.options or {}) else (None, f"not one of the options: {raw!r}")
|
|
56
|
+
return coerce_json_value(raw, f.type)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def value_from_argument(name: str, f: TExtractField, raw: Any, conf: float, tracker: Tracker) -> TFieldValue:
|
|
60
|
+
"""One argument of the call as a field value; the call's confidence applies to every field."""
|
|
61
|
+
if raw is None:
|
|
62
|
+
return TFieldValue(None, conf)
|
|
63
|
+
value, problem = _coerce(f, raw)
|
|
64
|
+
if problem is not None:
|
|
65
|
+
tracker.warn(problem, name)
|
|
66
|
+
return EMPTY
|
|
67
|
+
if value is None:
|
|
68
|
+
return TFieldValue(None, conf)
|
|
69
|
+
if f.type == "enum":
|
|
70
|
+
rest = round((1 - conf) / (len(f.options or {}) - 1), DIGITS)
|
|
71
|
+
probs = {k: (conf if k == value else rest) for k in f.options or {}}
|
|
72
|
+
return TFieldValue(value, conf, probabilities=probs)
|
|
73
|
+
if f.type == "boolean":
|
|
74
|
+
return TFieldValue(value, conf, probability=round(conf if value else 1 - conf, DIGITS))
|
|
75
|
+
return TFieldValue(value, conf)
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
|
|
@@ -0,0 +1,154 @@
|
|
|
1
|
+
"""Engine for any server with an OpenAI-style /chat/completions endpoint that
|
|
2
|
+
returns logprobs: Ollama (>= 0.12.11), vLLM, llama.cpp server, LM Studio, or a
|
|
3
|
+
hosted provider your company already approved.
|
|
4
|
+
|
|
5
|
+
The model is asked to reply with a single label, and the answer is read from
|
|
6
|
+
the probabilities the model gave each candidate token. One short
|
|
7
|
+
completion per question; no text generation, no JSON parsing.
|
|
8
|
+
|
|
9
|
+
It is also an ``IGenerator``: for extraction fields that need free text
|
|
10
|
+
(names, amounts, dates) it generates one JSON object, and returns the tokens'
|
|
11
|
+
probabilities so each value gets its own confidence.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import math
|
|
17
|
+
from typing import Optional
|
|
18
|
+
|
|
19
|
+
import httpx
|
|
20
|
+
|
|
21
|
+
from dev_double.core.decision.errors import EngineError
|
|
22
|
+
from dev_double.core.decision.i_engine import IEngine
|
|
23
|
+
from dev_double.core.decision.i_generator import IGenerator
|
|
24
|
+
from dev_double.core.decision.label_scoring import label_distribution, label_from_text, one_hot, uniform
|
|
25
|
+
from dev_double.core.decision.t_generate import TGenerateQuery, TGenerateResult
|
|
26
|
+
from dev_double.core.decision.t_label_query import TLabelQuery, TLabelResult, TPosition
|
|
27
|
+
from dev_double.core.decision.t_usage import TUsage
|
|
28
|
+
|
|
29
|
+
# Enough tokens to get past a leading newline and to tell apart options that
|
|
30
|
+
# share their first token ("re" + "fund" vs "re" + "booking").
|
|
31
|
+
MAX_TOKENS = 8
|
|
32
|
+
TOP_LOGPROBS = 20
|
|
33
|
+
LOW_MASS = 0.5
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class EngineOpenAIImpl(IEngine, IGenerator):
|
|
37
|
+
name = "openai"
|
|
38
|
+
|
|
39
|
+
def __init__(
|
|
40
|
+
self,
|
|
41
|
+
base_url: str,
|
|
42
|
+
model: str,
|
|
43
|
+
api_key: Optional[str] = None,
|
|
44
|
+
timeout: float = 120.0,
|
|
45
|
+
client: Optional[httpx.AsyncClient] = None,
|
|
46
|
+
) -> None:
|
|
47
|
+
self._model = model
|
|
48
|
+
self._url = base_url.rstrip("/") + "/chat/completions"
|
|
49
|
+
self._headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
|
|
50
|
+
self._client = client or httpx.AsyncClient(timeout=timeout)
|
|
51
|
+
|
|
52
|
+
@property
|
|
53
|
+
def model(self) -> str:
|
|
54
|
+
return self._model
|
|
55
|
+
|
|
56
|
+
async def aclose(self) -> None:
|
|
57
|
+
await self._client.aclose()
|
|
58
|
+
|
|
59
|
+
async def _complete(self, system: str, user: str, **options: object) -> tuple[dict, TUsage]:
|
|
60
|
+
"""One /chat/completions call; returns the first choice and the usage."""
|
|
61
|
+
body = {
|
|
62
|
+
"model": self._model,
|
|
63
|
+
"messages": [{"role": "system", "content": system}, {"role": "user", "content": user}],
|
|
64
|
+
"temperature": 0,
|
|
65
|
+
"logprobs": True,
|
|
66
|
+
**options,
|
|
67
|
+
}
|
|
68
|
+
try:
|
|
69
|
+
response = await self._client.post(self._url, json=body, headers=self._headers)
|
|
70
|
+
except httpx.HTTPError as exc:
|
|
71
|
+
raise EngineError(f"Could not reach model server at {self._url}: {exc}") from exc
|
|
72
|
+
if response.status_code >= 400:
|
|
73
|
+
raise EngineError(
|
|
74
|
+
f"Model server returned {response.status_code}: {response.text[:500]}"
|
|
75
|
+
)
|
|
76
|
+
|
|
77
|
+
data = response.json()
|
|
78
|
+
try:
|
|
79
|
+
choice = data["choices"][0]
|
|
80
|
+
except (KeyError, IndexError) as exc:
|
|
81
|
+
raise EngineError(f"Unexpected response from model server: {str(data)[:500]}") from exc
|
|
82
|
+
|
|
83
|
+
raw_usage = data.get("usage") or {}
|
|
84
|
+
usage = TUsage(
|
|
85
|
+
input_tokens=raw_usage.get("prompt_tokens", 0),
|
|
86
|
+
output_tokens=raw_usage.get("completion_tokens", 0),
|
|
87
|
+
)
|
|
88
|
+
return choice, usage
|
|
89
|
+
|
|
90
|
+
async def distribution(self, query: TLabelQuery) -> TLabelResult:
|
|
91
|
+
choice, usage = await self._complete(
|
|
92
|
+
query.system, query.user, max_tokens=MAX_TOKENS, top_logprobs=TOP_LOGPROBS
|
|
93
|
+
)
|
|
94
|
+
content = (choice.get("message") or {}).get("content") or ""
|
|
95
|
+
positions = ((choice.get("logprobs") or {}).get("content")) or []
|
|
96
|
+
return self._read(positions, content, query.labels, usage)
|
|
97
|
+
|
|
98
|
+
async def generate(self, query: TGenerateQuery) -> TGenerateResult:
|
|
99
|
+
"""Free text (or one JSON object) at temperature 0; confidence is exp(mean
|
|
100
|
+
logprob) of the generated tokens, which are returned for per-value scoring."""
|
|
101
|
+
options: dict = {"max_tokens": query.max_tokens, "top_logprobs": 1}
|
|
102
|
+
if query.json:
|
|
103
|
+
options["response_format"] = {"type": "json_object"}
|
|
104
|
+
choice, usage = await self._complete(query.system, query.user, **options)
|
|
105
|
+
content = (choice.get("message") or {}).get("content") or ""
|
|
106
|
+
positions = ((choice.get("logprobs") or {}).get("content")) or []
|
|
107
|
+
tokens = [(p.get("token", ""), float(p["logprob"])) for p in positions if "logprob" in p]
|
|
108
|
+
warnings: list[str] = []
|
|
109
|
+
if tokens:
|
|
110
|
+
confidence = min(1.0, math.exp(sum(lp for _, lp in tokens) / len(tokens)))
|
|
111
|
+
else:
|
|
112
|
+
confidence = 1.0
|
|
113
|
+
warnings.append(
|
|
114
|
+
"The model server returned no logprobs, so extracted values' confidence is a "
|
|
115
|
+
"placeholder 1.0. Use a server with logprobs support (e.g. Ollama >= 0.12.11)."
|
|
116
|
+
)
|
|
117
|
+
if choice.get("finish_reason") == "length":
|
|
118
|
+
warnings.append(f"The output was cut off at {query.max_tokens} tokens.")
|
|
119
|
+
return TGenerateResult(content, confidence, usage, warnings, tokens)
|
|
120
|
+
|
|
121
|
+
@staticmethod
|
|
122
|
+
def _read(raw_positions: list[dict], content: str, labels: list[str], usage: TUsage) -> TLabelResult:
|
|
123
|
+
warnings: list[str] = []
|
|
124
|
+
positions = [
|
|
125
|
+
TPosition(
|
|
126
|
+
token=pos.get("token", ""),
|
|
127
|
+
candidates=[(c["token"], c["logprob"]) for c in pos.get("top_logprobs") or []],
|
|
128
|
+
)
|
|
129
|
+
for pos in raw_positions
|
|
130
|
+
]
|
|
131
|
+
dist, mass = label_distribution(positions, labels)
|
|
132
|
+
if mass > 0:
|
|
133
|
+
if mass < LOW_MASS:
|
|
134
|
+
warnings.append(
|
|
135
|
+
f"Only {mass:.0%} of the model's probability landed on a valid label; "
|
|
136
|
+
"this answer is unreliable. Try a larger or instruction-tuned model."
|
|
137
|
+
)
|
|
138
|
+
return TLabelResult(dist, usage, warnings)
|
|
139
|
+
|
|
140
|
+
# No usable logprobs: fall back to the generated text.
|
|
141
|
+
if not positions:
|
|
142
|
+
warnings.append(
|
|
143
|
+
"The model server returned no logprobs, so probabilities are 0/1 guesses, "
|
|
144
|
+
"not calibrated. Use a server with logprobs support (e.g. Ollama >= 0.12.11)."
|
|
145
|
+
)
|
|
146
|
+
match = label_from_text(content, labels)
|
|
147
|
+
if match is not None:
|
|
148
|
+
return TLabelResult(one_hot(labels, match), usage, warnings)
|
|
149
|
+
warnings.append(
|
|
150
|
+
f"The model did not answer with a valid label (got {content[:40]!r}); "
|
|
151
|
+
"returning a uniform distribution. Reasoning models that emit <think> "
|
|
152
|
+
"first are not supported; use a non-thinking instruct model."
|
|
153
|
+
)
|
|
154
|
+
return TLabelResult(uniform(labels), usage, warnings)
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
|
|
@@ -0,0 +1,100 @@
|
|
|
1
|
+
"""A real System 1 decision model (Kev, Laya) served at ``POST /v1/systemone``.
|
|
2
|
+
|
|
3
|
+
Each question is sent as-is (no dev-double prompt): these models answer
|
|
4
|
+
typed questions natively, so this implements ``IDecider`` directly.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import asyncio
|
|
10
|
+
from typing import Any, Optional
|
|
11
|
+
|
|
12
|
+
import httpx
|
|
13
|
+
|
|
14
|
+
from dev_double.core.decision.answer_shaping import shape_answer
|
|
15
|
+
from dev_double.core.decision.errors import EngineError
|
|
16
|
+
from dev_double.core.decision.i_decider import IDecider
|
|
17
|
+
from dev_double.core.decision.t_answer import TAnswer
|
|
18
|
+
from dev_double.core.decision.t_input import TInputValue
|
|
19
|
+
from dev_double.core.decision.t_question import TBinaryQuestion, TChoiceQuestion, TQuestion
|
|
20
|
+
from dev_double.core.decision.t_usage import TUsage
|
|
21
|
+
from dev_double.core.decision.tracker import Tracker
|
|
22
|
+
|
|
23
|
+
PATH = "/v1/systemone"
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def to_wire(q: TQuestion) -> dict[str, Any]:
|
|
27
|
+
"""A core question in the System 1 wire format."""
|
|
28
|
+
if isinstance(q, TBinaryQuestion):
|
|
29
|
+
body: dict[str, Any] = {"type": "noul", "instructions": q.question}
|
|
30
|
+
criteria = {k: v for k, v in (("true", q.yes), ("false", q.no)) if v}
|
|
31
|
+
if criteria:
|
|
32
|
+
body["criteria"] = criteria
|
|
33
|
+
return body
|
|
34
|
+
if isinstance(q, TChoiceQuestion):
|
|
35
|
+
return {"type": "choice", "instructions": q.question, "criteria": q.options}
|
|
36
|
+
return {"type": "score", "instructions": q.question, "criteria": q.levels}
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def from_wire(q: TQuestion, answer: dict[str, Any]) -> dict[str, float]:
|
|
40
|
+
"""The wire answer as a distribution over the question's labels."""
|
|
41
|
+
if isinstance(q, TBinaryQuestion):
|
|
42
|
+
p = float(answer["noul"])
|
|
43
|
+
return {"yes": p, "no": 1 - p}
|
|
44
|
+
return {str(k): float(v) for k, v in answer["probabilities"].items()}
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class DeciderSystemOneImpl(IDecider):
|
|
48
|
+
def __init__(
|
|
49
|
+
self,
|
|
50
|
+
name: str,
|
|
51
|
+
base_url: str,
|
|
52
|
+
api_key: Optional[str] = None,
|
|
53
|
+
*,
|
|
54
|
+
model: Optional[str] = None,
|
|
55
|
+
max_concurrency: int = 1,
|
|
56
|
+
timeout: float = 120.0,
|
|
57
|
+
client: Optional[httpx.AsyncClient] = None,
|
|
58
|
+
) -> None:
|
|
59
|
+
self._name = name
|
|
60
|
+
self._model = model or name
|
|
61
|
+
self._base_url = base_url.rstrip("/")
|
|
62
|
+
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
|
|
63
|
+
self._client = client or httpx.AsyncClient(timeout=timeout)
|
|
64
|
+
self._headers = headers
|
|
65
|
+
self._limit = asyncio.Semaphore(max_concurrency)
|
|
66
|
+
|
|
67
|
+
@property
|
|
68
|
+
def name(self) -> str:
|
|
69
|
+
return self._name
|
|
70
|
+
|
|
71
|
+
@property
|
|
72
|
+
def model(self) -> str:
|
|
73
|
+
return self._model
|
|
74
|
+
|
|
75
|
+
async def ask(
|
|
76
|
+
self, value: TInputValue, question: TQuestion, tracker: Tracker, label: str = ""
|
|
77
|
+
) -> TAnswer:
|
|
78
|
+
payload = {"state": value, "questions": {"q": to_wire(question)}}
|
|
79
|
+
url = self._base_url + PATH
|
|
80
|
+
try:
|
|
81
|
+
async with self._limit:
|
|
82
|
+
r = await self._client.post(url, json=payload, headers=self._headers)
|
|
83
|
+
except httpx.HTTPError as exc:
|
|
84
|
+
raise EngineError(f"Could not reach System 1 server at {url}: {exc}") from exc
|
|
85
|
+
if r.status_code >= 400:
|
|
86
|
+
raise EngineError(f"System 1 server returned {r.status_code}: {r.text[:500]}")
|
|
87
|
+
data = r.json()
|
|
88
|
+
try:
|
|
89
|
+
probs = from_wire(question, data["answers"]["q"])
|
|
90
|
+
except (KeyError, TypeError, ValueError) as exc:
|
|
91
|
+
raise EngineError(f"Unexpected response from System 1 server: {str(data)[:500]}") from exc
|
|
92
|
+
usage = data.get("usage") or {}
|
|
93
|
+
tracker.usage.add(
|
|
94
|
+
TUsage(input_tokens=usage.get("input_tokens", 0), output_tokens=usage.get("output_tokens", 0))
|
|
95
|
+
)
|
|
96
|
+
# The server's probabilities are passed through unrounded, as it sent them.
|
|
97
|
+
return shape_answer(question, probs, round_probs=False)
|
|
98
|
+
|
|
99
|
+
async def aclose(self) -> None:
|
|
100
|
+
await self._client.aclose()
|