init the project
This commit is contained in:
@@ -0,0 +1,145 @@
|
||||
# OpenAI-compatible LLM client with retry, token tracking, and vision support.
|
||||
|
||||
import base64
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
from openai import OpenAI
|
||||
|
||||
logger = logging.getLogger("testflow")
|
||||
|
||||
|
||||
class LLMClient:
|
||||
"""Generic OpenAI-compatible LLM client.
|
||||
|
||||
Usage::
|
||||
|
||||
client = LLMClient(api_key="sk-xxx", base_url="https://api.deepseek.com/v1")
|
||||
text = client.chat("deepseek-chat", [{"role": "user", "content": "Hello"}])
|
||||
print(client.usage)
|
||||
"""
|
||||
|
||||
TIMEOUT = 120
|
||||
MAX_RETRIES = 3
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
api_key: str = "",
|
||||
base_url: str = "",
|
||||
model: str = "",
|
||||
timeout: int | None = None,
|
||||
):
|
||||
if not api_key:
|
||||
raise ValueError(f"API key is required for LLMClient")
|
||||
self._client = OpenAI(api_key=api_key, base_url=base_url)
|
||||
self._timeout = timeout or self.TIMEOUT
|
||||
self._model = model
|
||||
self._prompt_tokens = 0
|
||||
self._completion_tokens = 0
|
||||
|
||||
@property
|
||||
def usage(self) -> dict:
|
||||
return {
|
||||
"prompt_tokens": self._prompt_tokens,
|
||||
"completion_tokens": self._completion_tokens,
|
||||
"total_tokens": self._prompt_tokens + self._completion_tokens,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def estimate_tokens(text: str) -> int:
|
||||
cjk = sum(1 for c in text if '一' <= c <= '鿿' or ' ' <= c <= '〿')
|
||||
other = len(text) - cjk
|
||||
return max(1, int(cjk / 1.7 + other / 3.0))
|
||||
|
||||
@staticmethod
|
||||
def estimate_image_tokens() -> int:
|
||||
return 500
|
||||
|
||||
def chat(
|
||||
self,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
*,
|
||||
timeout: int | None = None,
|
||||
response_format: dict | None = None,
|
||||
temperature: float = 0.0,
|
||||
) -> str:
|
||||
# Estimate prompt size
|
||||
prompt_chars = sum(len(str(m.get("content", ""))) for m in messages)
|
||||
label = f"chat({model})"
|
||||
|
||||
def _call():
|
||||
t0 = time.time()
|
||||
kwargs = dict(
|
||||
model=model,
|
||||
messages=messages,
|
||||
timeout=timeout or self._timeout,
|
||||
temperature=temperature,
|
||||
)
|
||||
if response_format is not None:
|
||||
kwargs["response_format"] = response_format
|
||||
logger.info("[LLM] → %s (prompt ~%d chars / ~%d tokens)",
|
||||
model, prompt_chars, prompt_chars // 3)
|
||||
resp = self._client.chat.completions.create(**kwargs)
|
||||
content = resp.choices[0].message.content
|
||||
usg = resp.usage
|
||||
if usg:
|
||||
self._prompt_tokens += usg.prompt_tokens
|
||||
self._completion_tokens += usg.completion_tokens
|
||||
elapsed = time.time() - t0
|
||||
logger.info("[LLM] ← %s: %d chars, p:%d c:%d tokens, %.1fs",
|
||||
model, len(content) if content else 0,
|
||||
usg.prompt_tokens if usg else 0,
|
||||
usg.completion_tokens if usg else 0,
|
||||
elapsed)
|
||||
if not content:
|
||||
raise RuntimeError("Empty response from LLM")
|
||||
return content
|
||||
|
||||
return self._retry(_call, label)
|
||||
|
||||
def chat_with_image(
|
||||
self,
|
||||
model: str,
|
||||
image_path: str,
|
||||
prompt: str,
|
||||
*,
|
||||
timeout: int | None = None,
|
||||
) -> str:
|
||||
"""Send a prompt with an image attachment (for vision models)."""
|
||||
with open(image_path, "rb") as f:
|
||||
img_b64 = base64.b64encode(f.read()).decode()
|
||||
mime = self._mime_type(image_path)
|
||||
|
||||
messages = [{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": f"data:{mime};base64,{img_b64}"}},
|
||||
{"type": "text", "text": prompt},
|
||||
],
|
||||
}]
|
||||
return self.chat(model, messages, timeout=timeout)
|
||||
|
||||
@staticmethod
|
||||
def _mime_type(image_path: str) -> str:
|
||||
ext = os.path.splitext(image_path)[1].lstrip(".").lower()
|
||||
return {
|
||||
"png": "image/png", "jpg": "image/jpeg", "jpeg": "image/jpeg",
|
||||
"gif": "image/gif", "bmp": "image/bmp",
|
||||
"webp": "image/webp", "svg": "image/svg+xml", "tiff": "image/tiff",
|
||||
}.get(ext, "image/png")
|
||||
|
||||
def _retry(self, fn, label: str) -> str:
|
||||
last_error: Optional[Exception] = None
|
||||
for attempt in range(self.MAX_RETRIES):
|
||||
try:
|
||||
return fn()
|
||||
except Exception as e:
|
||||
last_error = e
|
||||
logger.warning("%s error (attempt %d/%d): %s", label, attempt + 1, self.MAX_RETRIES, e)
|
||||
if attempt < self.MAX_RETRIES - 1:
|
||||
time.sleep(2 ** attempt)
|
||||
raise RuntimeError(f"{label}: all retries exhausted") from last_error
|
||||
Reference in New Issue
Block a user