annator / src /services /llm_client.py
techprotrade's picture
Add src directory
734b5b4 verified
Raw
History Blame Contribute Delete
9.55 kB
from __future__ import annotations
from typing import Optional, Dict, Any, List
from dataclasses import dataclass
from pathlib import Path
import os
import requests
PROVIDER_DEFAULTS: dict[str, dict[str, str]] = {
"lmstudio": {
"base_url": "http://localhost:1234/v1",
"api_key": "not-needed",
"model": "local-model",
},
"ollama": {
"base_url": "http://localhost:11434/v1",
"api_key": "not-needed",
"model": "llama3",
},
"openwebui": {
"base_url": "http://localhost:8080/v1",
"api_key": "not-needed",
"model": "llama3",
},
"grok": {
"base_url": "https://api.x.ai/v1",
"api_key": "",
"model": "grok-2-latest",
},
"deepseek": {
"base_url": "https://api.deepseek.com/v1",
"api_key": "",
"model": "deepseek-chat",
},
"copilot": {
"base_url": "https://models.github.ai/inference",
"api_key": "",
"model": "openai/gpt-4.1-mini",
},
}
_ENV_CACHE: dict[str, str] | None = None
@dataclass
class LLMConfig:
base_url: str = "http://localhost:1234/v1"
api_key: str = "not-needed"
model: str = "local-model"
max_tokens: int = 2048
temperature: float = 0.7
timeout: int = 120
provider: str = "lmstudio"
def _clean_env_value(value: str) -> str:
cleaned = value.strip()
if len(cleaned) >= 2 and cleaned[0] == cleaned[-1] and cleaned[0] in {"'", '"'}:
return cleaned[1:-1]
return cleaned
def _read_env_file(env_values: dict[str, str], env_file: Path) -> None:
for raw_line in env_file.read_text(encoding="utf-8").splitlines():
line = raw_line.strip()
if not line or line.startswith("#") or "=" not in line:
continue
if line.startswith("export "):
line = line[len("export ") :]
key, value = line.split("=", 1)
key = key.strip()
if not key:
continue
env_values[key] = _clean_env_value(value)
def _normalize_base_url(url: str, ensure_v1: bool = True) -> str:
cleaned = url.rstrip("/")
if ensure_v1 and not cleaned.endswith("/v1"):
return f"{cleaned}/v1"
return cleaned
def _first_present(env: dict[str, str], keys: list[str]) -> str | None:
for key in keys:
value = env.get(key)
if value:
return value
return None
def load_project_env(project_root: Path | None = None) -> dict[str, str]:
global _ENV_CACHE
if _ENV_CACHE is not None:
return _ENV_CACHE
env_values = dict(os.environ)
shared_env = Path(r"C:\Luuna\.env")
if shared_env.exists():
_read_env_file(env_values, shared_env)
search_root = Path(project_root) if project_root else Path.cwd()
candidates = [search_root, *search_root.parents]
for base in candidates:
env_file = base / ".env"
if env_file.exists():
_read_env_file(env_values, env_file)
break
_ENV_CACHE = env_values
return env_values
def reset_env_cache() -> None:
global _ENV_CACHE
_ENV_CACHE = None
def _provider_defaults(provider: str) -> dict[str, str]:
normalized = provider.lower()
if normalized not in PROVIDER_DEFAULTS:
raise ValueError(
f"Unknown provider '{provider}'. "
f"Available providers: {', '.join(sorted(PROVIDER_DEFAULTS.keys()))}"
)
return PROVIDER_DEFAULTS[normalized]
def resolve_llm_config(
provider: str,
model: str | None = None,
base_url: str | None = None,
api_key: str | None = None,
project_root: Path | None = None,
) -> LLMConfig:
normalized = provider.lower()
defaults = _provider_defaults(normalized)
env = load_project_env(project_root)
prefix = normalized.upper()
base_url_aliases = {
"lmstudio": [f"{prefix}_BASE_URL", "LOCAL_API_URL"],
"ollama": [f"{prefix}_BASE_URL", "OLLAMA_HOST"],
"grok": [f"{prefix}_BASE_URL"],
"deepseek": [f"{prefix}_BASE_URL"],
"copilot": [f"{prefix}_BASE_URL"],
"openwebui": [f"{prefix}_BASE_URL"],
}
api_key_aliases = {
"lmstudio": [f"{prefix}_API_KEY"],
"ollama": [f"{prefix}_API_KEY", "OLLAMA_API_KEY"],
"grok": [f"{prefix}_API_KEY", "GROQ_API_KEY", "GROQ_API_KEY_ALT"],
"deepseek": [f"{prefix}_API_KEY"],
"copilot": [f"{prefix}_API_KEY", "GITHUB_TOKEN", "REFINED_GITHUB_TOKEN"],
"openwebui": [f"{prefix}_API_KEY"],
}
model_aliases = {
"lmstudio": [f"{prefix}_MODEL", "ACTIVE_MODEL"],
"ollama": [f"{prefix}_MODEL"],
"grok": [f"{prefix}_MODEL"],
"deepseek": [f"{prefix}_MODEL"],
"copilot": [f"{prefix}_MODEL"],
"openwebui": [f"{prefix}_MODEL"],
}
resolved_base_url = (
base_url
or _first_present(env, base_url_aliases.get(normalized, [f"{prefix}_BASE_URL"]))
or defaults["base_url"]
)
resolved_api_key = (
api_key
or _first_present(env, api_key_aliases.get(normalized, [f"{prefix}_API_KEY"]))
or defaults["api_key"]
)
resolved_model = (
model
or _first_present(env, model_aliases.get(normalized, [f"{prefix}_MODEL"]))
or defaults["model"]
)
if normalized in {"lmstudio", "ollama", "openwebui", "grok", "deepseek", "copilot"}:
resolved_base_url = _normalize_base_url(resolved_base_url)
return LLMConfig(
base_url=resolved_base_url,
api_key=resolved_api_key,
model=resolved_model,
provider=normalized,
)
class LLMClient:
def __init__(self, config: Optional[LLMConfig] = None, provider: str = "lmstudio"):
self.config = config or LLMConfig(provider=provider)
self.provider = self.config.provider or provider
self._session = requests.Session()
if self.config.api_key and self.config.api_key != "not-needed":
self._session.headers.update(
{"Authorization": f"Bearer {self.config.api_key}"}
)
def complete(self, prompt: str, system: str = "", **kwargs) -> str:
messages = []
if system:
messages.append({"role": "system", "content": system})
messages.append({"role": "user", "content": prompt})
payload = {
"model": self.config.model,
"messages": messages,
"max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
"temperature": kwargs.get("temperature", self.config.temperature),
}
try:
response = self._session.post(
f"{self.config.base_url}/chat/completions",
json=payload,
timeout=self.config.timeout,
)
response.raise_for_status()
data = response.json()
return data["choices"][0]["message"]["content"]
except Exception as e:
return f"Error: {str(e)}"
def chat(self, messages: List[Dict[str, str]], **kwargs) -> str:
payload = {
"model": self.config.model,
"messages": messages,
"max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
"temperature": kwargs.get("temperature", self.config.temperature),
}
try:
response = self._session.post(
f"{self.config.base_url}/chat/completions",
json=payload,
timeout=self.config.timeout,
)
response.raise_for_status()
data = response.json()
return data["choices"][0]["message"]["content"]
except Exception as e:
return f"Error: {str(e)}"
@staticmethod
def for_lmstudio(
model: str = "local-model", base_url: str = "http://localhost:1234/v1"
):
return LLMClient(
LLMConfig(
base_url=base_url,
model=model,
api_key="not-needed",
provider="lmstudio",
),
"lmstudio",
)
@staticmethod
def for_ollama(model: str = "llama3", base_url: str = "http://localhost:11434/v1"):
return LLMClient(
LLMConfig(
base_url=base_url,
model=model,
api_key="not-needed",
provider="ollama",
),
"ollama",
)
@staticmethod
def for_openwebui(
model: str = "llama3", base_url: str = "http://localhost:8080/v1"
):
return LLMClient(
LLMConfig(
base_url=base_url,
model=model,
api_key="not-needed",
provider="openwebui",
),
"openwebui",
)
@staticmethod
def for_provider(
provider: str,
model: str | None = None,
base_url: str | None = None,
api_key: str | None = None,
project_root: Path | None = None,
):
config = resolve_llm_config(
provider,
model=model,
base_url=base_url,
api_key=api_key,
project_root=project_root,
)
return LLMClient(config, provider=config.provider)