tryaii 0.3.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.
- tryaii/__init__.py +69 -0
- tryaii/async_client.py +488 -0
- tryaii/benchmarks/__init__.py +4 -0
- tryaii/benchmarks/registry.py +193 -0
- tryaii/benchmarks/standard.py +105 -0
- tryaii/budget.py +592 -0
- tryaii/cache/__init__.py +3 -0
- tryaii/cache/lru.py +76 -0
- tryaii/centroids/__init__.py +4 -0
- tryaii/centroids/data/centroids_all-MiniLM-L6-v2.json +1 -0
- tryaii/centroids/data/training_queries.json +246 -0
- tryaii/centroids/generator.py +214 -0
- tryaii/centroids/loader.py +214 -0
- tryaii/classifiers/__init__.py +8 -0
- tryaii/classifiers/base.py +80 -0
- tryaii/classifiers/embedding.py +290 -0
- tryaii/cli/__init__.py +0 -0
- tryaii/cli/banner.py +162 -0
- tryaii/cli/main.py +780 -0
- tryaii/client.py +207 -0
- tryaii/config.py +105 -0
- tryaii/embeddings/__init__.py +9 -0
- tryaii/embeddings/base.py +54 -0
- tryaii/embeddings/local.py +83 -0
- tryaii/embeddings/openai_provider.py +83 -0
- tryaii/integrations/__init__.py +3 -0
- tryaii/integrations/openrouter.py +441 -0
- tryaii/registry/__init__.py +3 -0
- tryaii/registry/models.py +252 -0
- tryaii/registry/presets/__init__.py +0 -0
- tryaii/registry/presets/default_models.json +932 -0
- tryaii/router.py +283 -0
- tryaii/scoring/__init__.py +11 -0
- tryaii/scoring/benchmarks.py +82 -0
- tryaii/scoring/engine.py +284 -0
- tryaii/scoring/priorities.py +93 -0
- tryaii-0.3.0.dist-info/METADATA +186 -0
- tryaii-0.3.0.dist-info/RECORD +41 -0
- tryaii-0.3.0.dist-info/WHEEL +4 -0
- tryaii-0.3.0.dist-info/entry_points.txt +2 -0
- tryaii-0.3.0.dist-info/licenses/LICENSE +190 -0
tryaii/__init__.py
ADDED
|
@@ -0,0 +1,69 @@
|
|
|
1
|
+
"""
|
|
2
|
+
TryAii-DRE -- Embedding-based AI Model Router
|
|
3
|
+
|
|
4
|
+
Understands your prompt semantically and routes to the best model
|
|
5
|
+
based on benchmarks, cost, speed, and quality priorities.
|
|
6
|
+
|
|
7
|
+
Usage:
|
|
8
|
+
from tryaii import Router
|
|
9
|
+
|
|
10
|
+
router = Router()
|
|
11
|
+
result = router.route("Write a Python function to merge sorted arrays")
|
|
12
|
+
print(result.best_model)
|
|
13
|
+
print(result.scores)
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
import logging
|
|
17
|
+
|
|
18
|
+
from tryaii.benchmarks.registry import BenchmarkRegistry
|
|
19
|
+
from tryaii.budget import (
|
|
20
|
+
DEFAULT_DIFFICULTY_GAMMA,
|
|
21
|
+
DEFAULT_DIFFICULTY_SOURCE,
|
|
22
|
+
BudgetCandidate,
|
|
23
|
+
BudgetedRouteResult,
|
|
24
|
+
BudgetOptimizationResult,
|
|
25
|
+
compute_difficulty,
|
|
26
|
+
estimate_tokens,
|
|
27
|
+
route_dataset_with_budget,
|
|
28
|
+
)
|
|
29
|
+
from tryaii.client import DREClient
|
|
30
|
+
from tryaii.config import TryaiiDreConfig
|
|
31
|
+
from tryaii.registry.models import ModelInfo, ModelRegistry
|
|
32
|
+
from tryaii.router import Router, RouteResult
|
|
33
|
+
from tryaii.scoring.priorities import DEFAULT_PRIORITIES, Priorities
|
|
34
|
+
|
|
35
|
+
# Attach a NullHandler so library logging stays silent unless the host app
|
|
36
|
+
# configures handlers. Done after imports to keep module-level imports at top.
|
|
37
|
+
logging.getLogger("tryaii").addHandler(logging.NullHandler())
|
|
38
|
+
|
|
39
|
+
__version__ = "0.3.0"
|
|
40
|
+
|
|
41
|
+
__all__ = [
|
|
42
|
+
"Router",
|
|
43
|
+
"RouteResult",
|
|
44
|
+
"ModelRegistry",
|
|
45
|
+
"ModelInfo",
|
|
46
|
+
"Priorities",
|
|
47
|
+
"DEFAULT_PRIORITIES",
|
|
48
|
+
"BenchmarkRegistry",
|
|
49
|
+
"TryaiiDreConfig",
|
|
50
|
+
"DREClient",
|
|
51
|
+
"AsyncDREClient",
|
|
52
|
+
"BudgetCandidate",
|
|
53
|
+
"BudgetOptimizationResult",
|
|
54
|
+
"BudgetedRouteResult",
|
|
55
|
+
"DEFAULT_DIFFICULTY_GAMMA",
|
|
56
|
+
"DEFAULT_DIFFICULTY_SOURCE",
|
|
57
|
+
"compute_difficulty",
|
|
58
|
+
"estimate_tokens",
|
|
59
|
+
"route_dataset_with_budget",
|
|
60
|
+
"__version__",
|
|
61
|
+
]
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def __getattr__(name: str):
|
|
65
|
+
if name == "AsyncDREClient":
|
|
66
|
+
from tryaii.async_client import AsyncDREClient
|
|
67
|
+
|
|
68
|
+
return AsyncDREClient
|
|
69
|
+
raise AttributeError(f"module 'tryaii' has no attribute {name!r}")
|
tryaii/async_client.py
ADDED
|
@@ -0,0 +1,488 @@
|
|
|
1
|
+
"""
|
|
2
|
+
AsyncDREClient -- async version of DREClient for TryAii-DRE.
|
|
3
|
+
|
|
4
|
+
Provides the same interface as DREClient but with async/await support.
|
|
5
|
+
Uses asyncio.to_thread() for CPU-bound Router calls and httpx.AsyncClient
|
|
6
|
+
for non-blocking API calls.
|
|
7
|
+
|
|
8
|
+
Usage:
|
|
9
|
+
from tryaii.async_client import AsyncDREClient
|
|
10
|
+
|
|
11
|
+
client = AsyncDREClient(api_key="sk-or-...")
|
|
12
|
+
response = await client.chat("Write a quicksort in Python")
|
|
13
|
+
print(response.model_used, response.content)
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import asyncio
|
|
19
|
+
import json
|
|
20
|
+
import logging
|
|
21
|
+
import math
|
|
22
|
+
import os
|
|
23
|
+
import random
|
|
24
|
+
from collections.abc import AsyncGenerator
|
|
25
|
+
from datetime import datetime, timezone
|
|
26
|
+
from email.utils import parsedate_to_datetime
|
|
27
|
+
from typing import Optional
|
|
28
|
+
|
|
29
|
+
from tryaii.classifiers.base import MAX_PROMPT_LENGTH
|
|
30
|
+
from tryaii.config import TryaiiDreConfig
|
|
31
|
+
from tryaii.integrations.openrouter import (
|
|
32
|
+
MODEL_ID_TO_OPENROUTER,
|
|
33
|
+
OpenRouterResponse,
|
|
34
|
+
)
|
|
35
|
+
from tryaii.router import Router, RouteResult
|
|
36
|
+
from tryaii.scoring.priorities import Priorities
|
|
37
|
+
|
|
38
|
+
logger = logging.getLogger("tryaii.async_client")
|
|
39
|
+
|
|
40
|
+
try:
|
|
41
|
+
import httpx
|
|
42
|
+
except ImportError: # pragma: no cover - exercised only without optional extra
|
|
43
|
+
httpx = None
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1"
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
class AsyncDREClient:
|
|
50
|
+
"""
|
|
51
|
+
Async high-level client that combines routing and API calls.
|
|
52
|
+
|
|
53
|
+
Same interface as DREClient but all methods are async. Router calls
|
|
54
|
+
are offloaded to a thread pool via asyncio.to_thread(), and API calls
|
|
55
|
+
use httpx.AsyncClient for true async I/O.
|
|
56
|
+
|
|
57
|
+
Args:
|
|
58
|
+
api_key: OpenRouter API key. Falls back to OPENROUTER_API_KEY env var.
|
|
59
|
+
priorities: Default priorities for all routing calls.
|
|
60
|
+
embedding_model: Sentence-transformers model name for embeddings.
|
|
61
|
+
"""
|
|
62
|
+
|
|
63
|
+
def __init__(
|
|
64
|
+
self,
|
|
65
|
+
api_key: Optional[str] = None,
|
|
66
|
+
priorities: Optional[Priorities] = None,
|
|
67
|
+
embedding_model: Optional[str] = None,
|
|
68
|
+
):
|
|
69
|
+
self._api_key = api_key or os.environ.get("OPENROUTER_API_KEY", "")
|
|
70
|
+
self._default_priorities = priorities
|
|
71
|
+
self._embedding_model = embedding_model
|
|
72
|
+
|
|
73
|
+
# Build config
|
|
74
|
+
config = TryaiiDreConfig()
|
|
75
|
+
if embedding_model:
|
|
76
|
+
config = TryaiiDreConfig(embedding_model=embedding_model)
|
|
77
|
+
|
|
78
|
+
# Core router (sync -- will be called via asyncio.to_thread)
|
|
79
|
+
self._router = Router(config=config)
|
|
80
|
+
|
|
81
|
+
# Async HTTP client (lazy-initialized)
|
|
82
|
+
self._http_client: Optional[httpx.AsyncClient] = None
|
|
83
|
+
|
|
84
|
+
def _ensure_http_client(self) -> httpx.AsyncClient:
|
|
85
|
+
"""Lazy-initialize the async HTTP client."""
|
|
86
|
+
if httpx is None:
|
|
87
|
+
raise ImportError(
|
|
88
|
+
"httpx is required for AsyncDREClient. "
|
|
89
|
+
"Install with: pip install tryaii[openrouter]"
|
|
90
|
+
)
|
|
91
|
+
if self._http_client is None:
|
|
92
|
+
self._http_client = httpx.AsyncClient(
|
|
93
|
+
base_url=OPENROUTER_BASE_URL,
|
|
94
|
+
headers={
|
|
95
|
+
"Authorization": f"Bearer {self._api_key}",
|
|
96
|
+
"X-Title": "tryaii",
|
|
97
|
+
"Content-Type": "application/json",
|
|
98
|
+
},
|
|
99
|
+
timeout=120.0,
|
|
100
|
+
)
|
|
101
|
+
return self._http_client
|
|
102
|
+
|
|
103
|
+
@staticmethod
|
|
104
|
+
def _resolve_model(model_id: str) -> str:
|
|
105
|
+
"""Convert TryAii-DRE model ID to OpenRouter slug."""
|
|
106
|
+
return MODEL_ID_TO_OPENROUTER.get(model_id, model_id)
|
|
107
|
+
|
|
108
|
+
# -- Validation / payload helpers ----------------------------------
|
|
109
|
+
|
|
110
|
+
@staticmethod
|
|
111
|
+
def _validate_prompt(prompt: str) -> str:
|
|
112
|
+
"""Reject empty/non-string prompts and truncate to MAX_PROMPT_LENGTH.
|
|
113
|
+
|
|
114
|
+
Mirrors the sync OpenRouterIntegration so both SDKs behave identically.
|
|
115
|
+
"""
|
|
116
|
+
if not prompt or not isinstance(prompt, str):
|
|
117
|
+
raise ValueError("prompt must be a non-empty string")
|
|
118
|
+
if len(prompt) > MAX_PROMPT_LENGTH:
|
|
119
|
+
prompt = prompt[:MAX_PROMPT_LENGTH]
|
|
120
|
+
return prompt
|
|
121
|
+
|
|
122
|
+
@staticmethod
|
|
123
|
+
def _build_payload(
|
|
124
|
+
openrouter_model: str,
|
|
125
|
+
messages: list,
|
|
126
|
+
temperature: float,
|
|
127
|
+
max_tokens: Optional[int],
|
|
128
|
+
stream: bool = False,
|
|
129
|
+
) -> dict:
|
|
130
|
+
"""Build an OpenRouter chat payload, mirroring the sync integration."""
|
|
131
|
+
payload: dict = {
|
|
132
|
+
"model": openrouter_model,
|
|
133
|
+
"messages": messages,
|
|
134
|
+
"temperature": temperature,
|
|
135
|
+
}
|
|
136
|
+
if stream:
|
|
137
|
+
payload["stream"] = True
|
|
138
|
+
# Only forward max_tokens when it is a positive, finite value.
|
|
139
|
+
if max_tokens is not None and math.isfinite(max_tokens) and max_tokens > 0:
|
|
140
|
+
payload["max_tokens"] = max_tokens
|
|
141
|
+
return payload
|
|
142
|
+
|
|
143
|
+
# -- Retry helper --------------------------------------------------
|
|
144
|
+
|
|
145
|
+
_RETRYABLE_STATUS_CODES = {429, 500, 502, 503, 504}
|
|
146
|
+
_MAX_RETRIES = 3
|
|
147
|
+
|
|
148
|
+
async def _post_with_retry(self, client, url: str, **kwargs):
|
|
149
|
+
"""POST with exponential-backoff retry on transient errors (async).
|
|
150
|
+
|
|
151
|
+
Retries on 429 (rate-limit) and 5xx server errors up to ``_MAX_RETRIES``
|
|
152
|
+
times, honoring the ``Retry-After`` header on 429. Mirrors the sync
|
|
153
|
+
``OpenRouterIntegration._request_with_retry``.
|
|
154
|
+
"""
|
|
155
|
+
last_exc: Exception | None = None
|
|
156
|
+
for attempt in range(self._MAX_RETRIES + 1):
|
|
157
|
+
try:
|
|
158
|
+
response = await client.post(url, **kwargs)
|
|
159
|
+
if response.status_code not in self._RETRYABLE_STATUS_CODES:
|
|
160
|
+
return response
|
|
161
|
+
last_exc = httpx.HTTPStatusError(
|
|
162
|
+
f"HTTP {response.status_code}",
|
|
163
|
+
request=response.request,
|
|
164
|
+
response=response,
|
|
165
|
+
)
|
|
166
|
+
if attempt == self._MAX_RETRIES:
|
|
167
|
+
response.raise_for_status()
|
|
168
|
+
return response # pragma: no cover
|
|
169
|
+
wait = self._backoff_wait(attempt, response)
|
|
170
|
+
except httpx.TransportError as exc:
|
|
171
|
+
last_exc = exc
|
|
172
|
+
if attempt == self._MAX_RETRIES:
|
|
173
|
+
raise
|
|
174
|
+
wait = 2 ** attempt + random.uniform(0, 1)
|
|
175
|
+
|
|
176
|
+
logger.warning(
|
|
177
|
+
"Async request to %s failed (attempt %d/%d), retrying in %.1fs",
|
|
178
|
+
url,
|
|
179
|
+
attempt + 1,
|
|
180
|
+
self._MAX_RETRIES + 1,
|
|
181
|
+
wait,
|
|
182
|
+
)
|
|
183
|
+
await asyncio.sleep(wait)
|
|
184
|
+
|
|
185
|
+
raise last_exc # type: ignore[misc]
|
|
186
|
+
|
|
187
|
+
@staticmethod
|
|
188
|
+
def _backoff_wait(attempt: int, response) -> float:
|
|
189
|
+
"""Calculate wait time, honoring Retry-After on 429."""
|
|
190
|
+
if response.status_code == 429:
|
|
191
|
+
retry_after = response.headers.get("Retry-After")
|
|
192
|
+
if retry_after is not None:
|
|
193
|
+
# Retry-After may be either delta-seconds or an HTTP-date.
|
|
194
|
+
try:
|
|
195
|
+
return float(retry_after)
|
|
196
|
+
except (ValueError, TypeError):
|
|
197
|
+
pass
|
|
198
|
+
try:
|
|
199
|
+
then = parsedate_to_datetime(retry_after)
|
|
200
|
+
if then is not None:
|
|
201
|
+
if then.tzinfo is None:
|
|
202
|
+
then = then.replace(tzinfo=timezone.utc)
|
|
203
|
+
delta = (then - datetime.now(timezone.utc)).total_seconds()
|
|
204
|
+
return max(0.0, delta)
|
|
205
|
+
except (ValueError, TypeError):
|
|
206
|
+
pass
|
|
207
|
+
return 2 ** attempt + random.uniform(0, 1)
|
|
208
|
+
|
|
209
|
+
@staticmethod
|
|
210
|
+
def _content_from_data(data: dict) -> str:
|
|
211
|
+
"""Extract message content, raising on a 200-with-error envelope.
|
|
212
|
+
|
|
213
|
+
OpenRouter may return HTTP 200 with an {"error": ...} body and no
|
|
214
|
+
choices; 'error' may be a dict or a bare string.
|
|
215
|
+
"""
|
|
216
|
+
choices = data.get("choices")
|
|
217
|
+
if not choices:
|
|
218
|
+
error = data.get("error")
|
|
219
|
+
if isinstance(error, dict):
|
|
220
|
+
msg = error.get("message", "No choices returned")
|
|
221
|
+
elif isinstance(error, str):
|
|
222
|
+
msg = error
|
|
223
|
+
else:
|
|
224
|
+
msg = "No choices returned"
|
|
225
|
+
raise ValueError(f"OpenRouter API error: {msg}")
|
|
226
|
+
return choices[0].get("message", {}).get("content", "")
|
|
227
|
+
|
|
228
|
+
@property
|
|
229
|
+
def router(self) -> Router:
|
|
230
|
+
"""Access the underlying Router instance."""
|
|
231
|
+
return self._router
|
|
232
|
+
|
|
233
|
+
async def route(
|
|
234
|
+
self,
|
|
235
|
+
prompt: str,
|
|
236
|
+
priorities: Optional[Priorities] = None,
|
|
237
|
+
top_k: int = 5,
|
|
238
|
+
) -> RouteResult:
|
|
239
|
+
"""
|
|
240
|
+
Route a prompt without making an API call (async).
|
|
241
|
+
|
|
242
|
+
The Router's classify/score logic is CPU-bound, so it is
|
|
243
|
+
offloaded to a thread to avoid blocking the event loop.
|
|
244
|
+
|
|
245
|
+
Args:
|
|
246
|
+
prompt: The user message to classify and route.
|
|
247
|
+
priorities: Override default priorities for this call.
|
|
248
|
+
top_k: Number of top models to include in results.
|
|
249
|
+
|
|
250
|
+
Returns:
|
|
251
|
+
RouteResult with best_model, scores, and classification.
|
|
252
|
+
"""
|
|
253
|
+
prio = priorities or self._default_priorities
|
|
254
|
+
return await asyncio.to_thread(
|
|
255
|
+
self._router.route, prompt, priorities=prio, top_k=top_k
|
|
256
|
+
)
|
|
257
|
+
|
|
258
|
+
async def chat(
|
|
259
|
+
self,
|
|
260
|
+
prompt: str,
|
|
261
|
+
priorities: Optional[Priorities] = None,
|
|
262
|
+
system_message: Optional[str] = None,
|
|
263
|
+
temperature: float = 0.7,
|
|
264
|
+
max_tokens: Optional[int] = None,
|
|
265
|
+
) -> OpenRouterResponse:
|
|
266
|
+
"""
|
|
267
|
+
Route the prompt to the best model and return the AI response (async).
|
|
268
|
+
|
|
269
|
+
Args:
|
|
270
|
+
prompt: The user message to route and send.
|
|
271
|
+
priorities: Override default priorities for this call.
|
|
272
|
+
system_message: Optional system prompt.
|
|
273
|
+
temperature: Sampling temperature (0.0 to 2.0).
|
|
274
|
+
max_tokens: Maximum tokens in the response.
|
|
275
|
+
|
|
276
|
+
Returns:
|
|
277
|
+
OpenRouterResponse with content, model_used, and routing info.
|
|
278
|
+
"""
|
|
279
|
+
prompt = self._validate_prompt(prompt)
|
|
280
|
+
|
|
281
|
+
# Route in a thread (CPU-bound)
|
|
282
|
+
route_result = await self.route(prompt, priorities=priorities)
|
|
283
|
+
model_id = route_result.best_model
|
|
284
|
+
reasoning = route_result.scores[0].reasoning if route_result.scores else ""
|
|
285
|
+
|
|
286
|
+
# Never POST an empty model -- the router could not score any model.
|
|
287
|
+
if not model_id:
|
|
288
|
+
raise ValueError("routing returned no model for this prompt")
|
|
289
|
+
|
|
290
|
+
openrouter_model = self._resolve_model(model_id)
|
|
291
|
+
|
|
292
|
+
# Build messages
|
|
293
|
+
messages = []
|
|
294
|
+
if system_message:
|
|
295
|
+
messages.append({"role": "system", "content": system_message})
|
|
296
|
+
messages.append({"role": "user", "content": prompt})
|
|
297
|
+
|
|
298
|
+
payload = self._build_payload(openrouter_model, messages, temperature, max_tokens)
|
|
299
|
+
|
|
300
|
+
# Async API call
|
|
301
|
+
client = self._ensure_http_client()
|
|
302
|
+
response = await self._post_with_retry(client, "/chat/completions", json=payload)
|
|
303
|
+
response.raise_for_status()
|
|
304
|
+
data = response.json()
|
|
305
|
+
|
|
306
|
+
content = self._content_from_data(data)
|
|
307
|
+
usage = data.get("usage", {})
|
|
308
|
+
|
|
309
|
+
return OpenRouterResponse(
|
|
310
|
+
content=content,
|
|
311
|
+
model_used=model_id,
|
|
312
|
+
openrouter_model=openrouter_model,
|
|
313
|
+
route_reasoning=reasoning,
|
|
314
|
+
usage=usage,
|
|
315
|
+
raw_response=data,
|
|
316
|
+
)
|
|
317
|
+
|
|
318
|
+
async def stream(
|
|
319
|
+
self,
|
|
320
|
+
prompt: str,
|
|
321
|
+
priorities: Optional[Priorities] = None,
|
|
322
|
+
system_message: Optional[str] = None,
|
|
323
|
+
temperature: float = 0.7,
|
|
324
|
+
max_tokens: Optional[int] = None,
|
|
325
|
+
) -> AsyncGenerator[str, None]:
|
|
326
|
+
"""
|
|
327
|
+
Route the prompt to the best model and stream the response (async).
|
|
328
|
+
|
|
329
|
+
Yields content chunks as they arrive from the API.
|
|
330
|
+
|
|
331
|
+
Args:
|
|
332
|
+
prompt: The user message to route and send.
|
|
333
|
+
priorities: Override default priorities for this call.
|
|
334
|
+
system_message: Optional system prompt.
|
|
335
|
+
temperature: Sampling temperature (0.0 to 2.0).
|
|
336
|
+
max_tokens: Maximum tokens in the response.
|
|
337
|
+
|
|
338
|
+
Yields:
|
|
339
|
+
String chunks of the response content.
|
|
340
|
+
"""
|
|
341
|
+
prompt = self._validate_prompt(prompt)
|
|
342
|
+
|
|
343
|
+
# Route in a thread (CPU-bound)
|
|
344
|
+
route_result = await self.route(prompt, priorities=priorities)
|
|
345
|
+
model_id = route_result.best_model
|
|
346
|
+
|
|
347
|
+
# Never POST an empty model -- the router could not score any model.
|
|
348
|
+
if not model_id:
|
|
349
|
+
raise ValueError("routing returned no model for this prompt")
|
|
350
|
+
|
|
351
|
+
openrouter_model = self._resolve_model(model_id)
|
|
352
|
+
|
|
353
|
+
# Build messages
|
|
354
|
+
messages = []
|
|
355
|
+
if system_message:
|
|
356
|
+
messages.append({"role": "system", "content": system_message})
|
|
357
|
+
messages.append({"role": "user", "content": prompt})
|
|
358
|
+
|
|
359
|
+
payload = self._build_payload(
|
|
360
|
+
openrouter_model, messages, temperature, max_tokens, stream=True
|
|
361
|
+
)
|
|
362
|
+
|
|
363
|
+
# Async streaming with reconnect on transient errors. Retries only cover
|
|
364
|
+
# the pre-first-byte connection phase: once a chunk has been yielded a
|
|
365
|
+
# mid-stream failure must re-raise rather than replay (duplicating
|
|
366
|
+
# already-emitted content).
|
|
367
|
+
client = self._ensure_http_client()
|
|
368
|
+
yielded_any = False
|
|
369
|
+
last_exc: Exception | None = None
|
|
370
|
+
for attempt in range(self._MAX_RETRIES + 1):
|
|
371
|
+
try:
|
|
372
|
+
async with client.stream(
|
|
373
|
+
"POST", "/chat/completions", json=payload
|
|
374
|
+
) as response:
|
|
375
|
+
response.raise_for_status()
|
|
376
|
+
async for line in response.aiter_lines():
|
|
377
|
+
if not line or not line.startswith("data: "):
|
|
378
|
+
continue
|
|
379
|
+
data_str = line[6:]
|
|
380
|
+
if data_str.strip() == "[DONE]":
|
|
381
|
+
break
|
|
382
|
+
try:
|
|
383
|
+
chunk = json.loads(data_str)
|
|
384
|
+
if "error" in chunk:
|
|
385
|
+
err_msg = chunk["error"]
|
|
386
|
+
if isinstance(err_msg, dict):
|
|
387
|
+
err_msg = err_msg.get("message", str(err_msg))
|
|
388
|
+
logger.error("Stream error from API: %s", err_msg)
|
|
389
|
+
raise ValueError(
|
|
390
|
+
f"OpenRouter stream error: {err_msg}"
|
|
391
|
+
)
|
|
392
|
+
delta = chunk["choices"][0].get("delta", {})
|
|
393
|
+
content = delta.get("content", "")
|
|
394
|
+
if content:
|
|
395
|
+
yielded_any = True
|
|
396
|
+
yield content
|
|
397
|
+
except (json.JSONDecodeError, KeyError, IndexError):
|
|
398
|
+
continue
|
|
399
|
+
return # Stream completed successfully; exit retry loop
|
|
400
|
+
except (httpx.TransportError, httpx.HTTPStatusError) as exc:
|
|
401
|
+
last_exc = exc
|
|
402
|
+
# Once content has started flowing, never retry from scratch.
|
|
403
|
+
if yielded_any or attempt == self._MAX_RETRIES:
|
|
404
|
+
raise
|
|
405
|
+
wait = 2 ** attempt + random.uniform(0, 1)
|
|
406
|
+
logger.warning(
|
|
407
|
+
"Async stream connection failed (attempt %d/%d), retrying in %.1fs",
|
|
408
|
+
attempt + 1,
|
|
409
|
+
self._MAX_RETRIES + 1,
|
|
410
|
+
wait,
|
|
411
|
+
)
|
|
412
|
+
await asyncio.sleep(wait)
|
|
413
|
+
|
|
414
|
+
raise last_exc # type: ignore[misc] # pragma: no cover
|
|
415
|
+
|
|
416
|
+
async def route_and_chat(
|
|
417
|
+
self,
|
|
418
|
+
prompt: str,
|
|
419
|
+
priorities: Optional[Priorities] = None,
|
|
420
|
+
system_message: Optional[str] = None,
|
|
421
|
+
temperature: float = 0.7,
|
|
422
|
+
max_tokens: Optional[int] = None,
|
|
423
|
+
) -> tuple[RouteResult, OpenRouterResponse]:
|
|
424
|
+
"""
|
|
425
|
+
Route a prompt and get the AI response, returning both (async).
|
|
426
|
+
|
|
427
|
+
Args:
|
|
428
|
+
prompt: The user message to route and send.
|
|
429
|
+
priorities: Override default priorities for this call.
|
|
430
|
+
system_message: Optional system prompt.
|
|
431
|
+
temperature: Sampling temperature (0.0 to 2.0).
|
|
432
|
+
max_tokens: Maximum tokens in the response.
|
|
433
|
+
|
|
434
|
+
Returns:
|
|
435
|
+
Tuple of (RouteResult, OpenRouterResponse).
|
|
436
|
+
"""
|
|
437
|
+
prompt = self._validate_prompt(prompt)
|
|
438
|
+
|
|
439
|
+
# Route first
|
|
440
|
+
route_result = await self.route(prompt, priorities=priorities)
|
|
441
|
+
model_id = route_result.best_model
|
|
442
|
+
reasoning = route_result.scores[0].reasoning if route_result.scores else ""
|
|
443
|
+
|
|
444
|
+
# Never POST an empty model -- the router could not score any model.
|
|
445
|
+
if not model_id:
|
|
446
|
+
raise ValueError("routing returned no model for this prompt")
|
|
447
|
+
|
|
448
|
+
openrouter_model = self._resolve_model(model_id)
|
|
449
|
+
|
|
450
|
+
# Build messages
|
|
451
|
+
messages = []
|
|
452
|
+
if system_message:
|
|
453
|
+
messages.append({"role": "system", "content": system_message})
|
|
454
|
+
messages.append({"role": "user", "content": prompt})
|
|
455
|
+
|
|
456
|
+
payload = self._build_payload(openrouter_model, messages, temperature, max_tokens)
|
|
457
|
+
|
|
458
|
+
# Async API call
|
|
459
|
+
client = self._ensure_http_client()
|
|
460
|
+
response = await self._post_with_retry(client, "/chat/completions", json=payload)
|
|
461
|
+
response.raise_for_status()
|
|
462
|
+
data = response.json()
|
|
463
|
+
|
|
464
|
+
content = self._content_from_data(data)
|
|
465
|
+
usage = data.get("usage", {})
|
|
466
|
+
|
|
467
|
+
api_response = OpenRouterResponse(
|
|
468
|
+
content=content,
|
|
469
|
+
model_used=model_id,
|
|
470
|
+
openrouter_model=openrouter_model,
|
|
471
|
+
route_reasoning=reasoning,
|
|
472
|
+
usage=usage,
|
|
473
|
+
raw_response=data,
|
|
474
|
+
)
|
|
475
|
+
|
|
476
|
+
return route_result, api_response
|
|
477
|
+
|
|
478
|
+
async def close(self) -> None:
|
|
479
|
+
"""Close the underlying HTTP client."""
|
|
480
|
+
if self._http_client is not None:
|
|
481
|
+
await self._http_client.aclose()
|
|
482
|
+
self._http_client = None
|
|
483
|
+
|
|
484
|
+
async def __aenter__(self) -> AsyncDREClient:
|
|
485
|
+
return self
|
|
486
|
+
|
|
487
|
+
async def __aexit__(self, exc_type, exc_val, exc_tb) -> None:
|
|
488
|
+
await self.close()
|