sdm-learn 0.1.0.dev0__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.
sdm/gateway.py ADDED
@@ -0,0 +1,169 @@
1
+ """The gateway to an LLM, and the environment checks that come with it.
2
+
3
+ Deliberately not called ``model``: in this project the skill document *is* the
4
+ model. This file is only the wire to a language model that reads it.
5
+
6
+ The bottom of the stack, so nothing here knows about skills, datasets, or the
7
+ learning loop. It turns a list of content blocks into one completion, retries
8
+ what is worth retrying, and refuses to paper over a truncated response. It is
9
+ also the one seam the offline checks replace, which is what keeps them honest:
10
+ there is no second path to the network.
11
+
12
+ ``SetupNeeded`` lives here because this is the module every other one can
13
+ reach. It marks the failures a traceback cannot fix: no API key, no SVG
14
+ rasterizer, no numpy.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ import base64
20
+ import contextlib
21
+ import io
22
+ import json
23
+ import os
24
+ import time
25
+ from pathlib import Path
26
+
27
+ import httpx
28
+
29
+
30
+ class SetupNeeded(Exception):
31
+ """Something about the environment needs a human, not a traceback."""
32
+
33
+
34
+ GATEWAY_URL = os.environ.get(
35
+ "SDM_GATEWAY_URL", "https://inference-api.nvidia.com/v1/chat/completions"
36
+ )
37
+ MODEL = os.environ.get("SDM_MODEL", "azure/anthropic/claude-opus-5")
38
+ EFFORT = os.environ.get("SDM_EFFORT", "low")
39
+ KEY_FILE = Path.home() / ".sdm_key"
40
+
41
+
42
+ def _api_key() -> str:
43
+ key = os.environ.get("SDM_API_KEY")
44
+ if key:
45
+ return key.strip()
46
+ path = Path(os.environ.get("SDM_API_KEY_FILE", KEY_FILE))
47
+ if path.is_file():
48
+ return path.read_text().strip()
49
+ raise SetupNeeded(
50
+ "no model to talk to yet. Set SDM_API_KEY, or write the key to "
51
+ f"{KEY_FILE}."
52
+ )
53
+
54
+
55
+ def client(read_timeout: float = 600.0, connect_timeout: float = 30.0):
56
+ """An authenticated session for the gateway."""
57
+ return httpx.Client(
58
+ headers={"Authorization": f"Bearer {_api_key()}",
59
+ "Content-Type": "application/json"},
60
+ timeout=httpx.Timeout(read_timeout, connect=connect_timeout),
61
+ )
62
+
63
+
64
+ @contextlib.contextmanager
65
+ def _borrowed(session):
66
+ """Yield a session, closing it only if we were the one who opened it."""
67
+ owned = session is None
68
+ session = session or client()
69
+ try:
70
+ yield session
71
+ finally:
72
+ if owned:
73
+ session.close()
74
+
75
+
76
+ def _stream(session, payload: dict, deadline: float = 1800.0):
77
+ chunks, finish = [], None
78
+ started = time.monotonic()
79
+ with session.stream("POST", GATEWAY_URL,
80
+ json={**payload, "stream": True}) as response:
81
+ if response.status_code != 200:
82
+ response.read()
83
+ raise httpx.HTTPStatusError(
84
+ f"{response.status_code}: {response.text[:200]}",
85
+ request=response.request, response=response,
86
+ )
87
+ for line in response.iter_lines():
88
+ if time.monotonic() - started > deadline:
89
+ raise httpx.TransportError(f"deadline {deadline}s hit")
90
+ if not line.startswith("data: "):
91
+ continue
92
+ body = line[6:]
93
+ if body.strip() == "[DONE]":
94
+ break
95
+ event = json.loads(body)
96
+ if not event.get("choices"):
97
+ continue
98
+ choice = event["choices"][0]
99
+ if (choice.get("delta") or {}).get("content"):
100
+ chunks.append(choice["delta"]["content"])
101
+ if choice.get("finish_reason"):
102
+ finish = choice["finish_reason"]
103
+ return "".join(chunks), finish
104
+
105
+
106
+ def complete(session, messages: list, max_tokens: int,
107
+ effort: str | None = None) -> str:
108
+ """One chat completion, retrying through the failures a long run hits."""
109
+ payload = {
110
+ "model": MODEL,
111
+ "max_tokens": max_tokens,
112
+ "output_config": {"effort": effort or EFFORT},
113
+ "messages": messages,
114
+ }
115
+ for attempt in range(30):
116
+ try:
117
+ content, finish = _stream(session, payload)
118
+ if finish == "content_filter":
119
+ raise RuntimeError("model refused to answer")
120
+ if not content or finish == "length":
121
+ payload["max_tokens"] = min(payload["max_tokens"] * 2, 32000)
122
+ print(f" [truncated ({finish}), retrying with "
123
+ f"max_tokens={payload['max_tokens']}]", flush=True)
124
+ continue
125
+ if finish != "stop":
126
+ print(f" [incomplete stream (finish={finish!r}), retrying]",
127
+ flush=True)
128
+ continue
129
+ return content
130
+ except (httpx.HTTPStatusError, httpx.TransportError,
131
+ json.JSONDecodeError) as error:
132
+ wait = min(15 * (attempt + 1), 600)
133
+ print(f" [retry in {wait}s: {type(error).__name__}: "
134
+ f"{str(error)[:120]}]", flush=True)
135
+ time.sleep(wait)
136
+ raise RuntimeError("too many retries")
137
+
138
+
139
+ def _png_block(png: bytes) -> dict:
140
+ """One content block holding already-encoded PNG bytes."""
141
+ encoded = base64.b64encode(png).decode("ascii")
142
+ return {"type": "image_url",
143
+ "image_url": {"url": f"data:image/png;base64,{encoded}"}}
144
+
145
+
146
+ def image_block(image, size: int = 224) -> dict:
147
+ """Encode an RGB array, PIL image, or image path as one content block."""
148
+ from PIL import Image
149
+
150
+ if isinstance(image, Image.Image):
151
+ pil = image.convert("RGB")
152
+ elif isinstance(image, (str, Path)):
153
+ with Image.open(image) as source:
154
+ pil = source.convert("RGB")
155
+ elif hasattr(image, "shape"):
156
+ import numpy as np
157
+
158
+ if image.ndim != 3 or image.shape[2] != 3:
159
+ raise ValueError("expected one HxWx3 RGB image")
160
+ pil = Image.fromarray(np.asarray(image).astype("uint8"), mode="RGB")
161
+ else:
162
+ raise TypeError("expected an RGB array, PIL image, or image path")
163
+ pil = pil.resize((size, size), Image.Resampling.NEAREST)
164
+ buffer = io.BytesIO()
165
+ pil.save(buffer, format="PNG", optimize=False)
166
+ return _png_block(buffer.getvalue())
167
+
168
+
169
+ # --------------------------------------------------------------------------