refactor: centralize llm provider transport
This commit is contained in:
@@ -0,0 +1,148 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from .stream import OpenAIStreamAccumulator
|
||||
|
||||
|
||||
class OpenAITransportError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class OpenAIHTTPError(OpenAITransportError):
|
||||
def __init__(self, code: int, detail: str = "") -> None:
|
||||
super().__init__(f"HTTP {code}")
|
||||
self.code = code
|
||||
self.detail = detail
|
||||
|
||||
def describe(self, label: str) -> str:
|
||||
suffix = f":{self.detail[:300]}" if self.detail else ""
|
||||
return f"{label}(HTTP {self.code}){suffix}"
|
||||
|
||||
|
||||
class OpenAIEmptyResponseError(OpenAITransportError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class OpenAIChatCompletion:
|
||||
content: Any
|
||||
latency_ms: int
|
||||
|
||||
|
||||
def chat_completion(
|
||||
*,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
messages: list[dict[str, Any]],
|
||||
timeout: int,
|
||||
user_agent: str,
|
||||
) -> OpenAIChatCompletion:
|
||||
request = _request(api_key, base_url, model, messages, user_agent, stream=False)
|
||||
started = time.perf_counter()
|
||||
try:
|
||||
with urllib.request.urlopen(request, timeout=timeout) as response:
|
||||
result = json.loads(response.read().decode("utf-8"))
|
||||
content = result["choices"][0]["message"]["content"]
|
||||
except urllib.error.HTTPError as exc:
|
||||
raise OpenAIHTTPError(exc.code, _http_error_detail(exc)) from exc
|
||||
except (
|
||||
urllib.error.URLError,
|
||||
TimeoutError,
|
||||
json.JSONDecodeError,
|
||||
KeyError,
|
||||
IndexError,
|
||||
) as exc:
|
||||
raise OpenAITransportError(str(exc)) from exc
|
||||
return OpenAIChatCompletion(
|
||||
content=content,
|
||||
latency_ms=round((time.perf_counter() - started) * 1000),
|
||||
)
|
||||
|
||||
|
||||
def stream_chat_completion(
|
||||
*,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
messages: list[dict[str, Any]],
|
||||
timeout: int,
|
||||
user_agent: str,
|
||||
) -> Iterator[str]:
|
||||
request = _request(api_key, base_url, model, messages, user_agent, stream=True)
|
||||
yielded = False
|
||||
try:
|
||||
with urllib.request.urlopen(request, timeout=timeout) as response:
|
||||
accumulator = OpenAIStreamAccumulator()
|
||||
for raw_line in response:
|
||||
line = raw_line.decode("utf-8", errors="replace").strip()
|
||||
if not line or line.startswith(":"):
|
||||
continue
|
||||
if line.startswith("data:"):
|
||||
line = line[5:].strip()
|
||||
if line == "[DONE]":
|
||||
break
|
||||
try:
|
||||
result = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
choices = result.get("choices") or []
|
||||
if not choices:
|
||||
continue
|
||||
content = accumulator.feed(choices[0] or {})
|
||||
if content:
|
||||
yielded = True
|
||||
yield str(content)
|
||||
except urllib.error.HTTPError as exc:
|
||||
raise OpenAIHTTPError(exc.code, _http_error_detail(exc)) from exc
|
||||
except (urllib.error.URLError, TimeoutError, OSError) as exc:
|
||||
raise OpenAITransportError(str(exc)) from exc
|
||||
if not yielded:
|
||||
raise OpenAIEmptyResponseError("empty response")
|
||||
|
||||
|
||||
def _request(
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
messages: list[dict[str, Any]],
|
||||
user_agent: str,
|
||||
*,
|
||||
stream: bool,
|
||||
) -> urllib.request.Request:
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"User-Agent": user_agent,
|
||||
}
|
||||
if stream:
|
||||
headers["Accept"] = "text/event-stream"
|
||||
return urllib.request.Request(
|
||||
f"{base_url.rstrip('/')}/chat/completions",
|
||||
data=json.dumps(
|
||||
{"model": model, "messages": messages, "stream": stream},
|
||||
ensure_ascii=False,
|
||||
).encode("utf-8"),
|
||||
headers=headers,
|
||||
method="POST",
|
||||
)
|
||||
|
||||
|
||||
def _http_error_detail(exc: urllib.error.HTTPError) -> str:
|
||||
try:
|
||||
payload = json.loads(exc.read().decode("utf-8", errors="replace"))
|
||||
error = payload.get("error")
|
||||
if isinstance(error, dict):
|
||||
return str(error.get("message") or error.get("code") or "")
|
||||
if error:
|
||||
return str(error)
|
||||
return str(payload.get("message") or "")
|
||||
except (json.JSONDecodeError, OSError):
|
||||
return ""
|
||||
Reference in New Issue
Block a user