Team Ai
Apppublic

openenv/echo_env

sourceHugging Faceupdated 1d agoView on Hugging Face
6likes
llm_client.py639 linesDownload Raw Back to core
1# SPDX-License-Identifier: BSD-3-Clause2 3"""LLM client abstraction for calling LLM endpoints.4 5Provides a generic RPC abstraction: point it at an endpoint/port, tell it the6protocol, and it works. OpenAI-compatible API is the first implementation,7covering OpenAI, vLLM, TGI, Ollama, HuggingFace Inference API, etc.8Anthropic's native API is supported via `AnthropicClient`.9 10Examples:11 12    ```python13    client = OpenAIClient("http://localhost", 8000, model="meta-llama/...")14    response = await client.complete("What is 2+2?")15 16    # The endpoint may carry its own port and path prefix:17    client = OpenAIClient("http://localhost:8000/v1", port=None, model="meta-llama/...")18 19    # Or use the factory for hosted APIs:20    client = create_llm_client("openai", model="gpt-4", api_key="sk-...")21    response = await client.complete_with_tools(messages, tools)22    ```23"""24 25from __future__ import annotations26 27import json28from abc import ABC, abstractmethod29from dataclasses import dataclass, field30from typing import Any31from urllib.parse import urlsplit, urlunsplit32 33from openai import AsyncOpenAI34 35_OPENAI_API_PREFIX = "/v1"36_MIN_PORT, _MAX_PORT = 1, 6553537 38 39def _redact_userinfo(endpoint: str) -> str:40    """Drop `user:password@` from a URL so it is safe to echo in an error."""41    scheme, sep, rest = endpoint.partition("://")42    authority, slash, tail = rest.partition("/")43    if not sep or "@" not in authority:44        return endpoint45    return f"{scheme}{sep}***@{authority.rsplit('@', 1)[-1]}{slash}{tail}"46 47 48def _join_endpoint_port(endpoint: str, port: int | None) -> str:49    """Validate an endpoint URL and combine it with an optional port.50 51    Only `http`/`https` URLs with a host are accepted. Credentials, query52    strings and fragments are rejected: the endpoint is logged and persisted in53    rollout metadata, and the SDKs build request URLs by appending to the path,54    which corrupts a query string. `port` is appended only when the URL does55    not name one; an explicit port that differs from the one in the URL is56    rejected rather than silently overridden. A trailing slash on the path is57    dropped. Invalid endpoints raise `ValueError`.58    """59    safe = _redact_userinfo(endpoint)60    try:61        parts = urlsplit(endpoint)62        if parts.scheme not in ("http", "https"):63            raise ValueError("expected an http:// or https:// URL")64        if not parts.hostname:65            raise ValueError("missing host")66        if parts.username is not None or parts.password is not None:67            raise ValueError("credentials in the URL are not supported, use api_key")68        if parts.query or parts.fragment:69            raise ValueError("query strings and fragments are not supported")70        if parts.netloc.endswith(":"):71            raise ValueError("empty port")72        url_port = parts.port73        for candidate in (url_port, port):74            if candidate is not None and not _MIN_PORT <= candidate <= _MAX_PORT:75                raise ValueError(76                    f"port {candidate} is out of range {_MIN_PORT}-{_MAX_PORT}"77                )78    except ValueError as exc:79        raise ValueError(f"Invalid endpoint URL {safe!r}: {exc}") from exc80 81    if url_port is None:82        if port is not None:83            parts = parts._replace(netloc=f"{parts.netloc}:{port}")84    elif port is not None and port != url_port:85        raise ValueError(86            f"Endpoint URL {safe!r} already specifies port {url_port}, "87            f"which conflicts with port={port}"88        )89    return urlunsplit(parts._replace(path=parts.path.rstrip("/")))90 91 92def _openai_base_url(base_url: str) -> str:93    """Append the OpenAI `/v1` prefix when `base_url` has no path.94 95    A URL that already names a path (`/v1`, or a gateway prefix such as96    `/openai/v1`) is used as-is, matching the OpenAI SDK convention that97    `base_url` includes the API prefix.98    """99    parts = urlsplit(base_url)100    if parts.path in ("", "/"):101        return urlunsplit(parts._replace(path=_OPENAI_API_PREFIX))102    return base_url103 104 105@dataclass106class ToolCall:107    """A single tool/function call returned by the model."""108 109    id: str110    name: str111    args: dict[str, Any]112 113 114@dataclass115class LLMResponse:116    """Normalized response from an LLM, with optional tool calls."""117 118    content: str119    tool_calls: list[ToolCall] = field(default_factory=list)120 121    def to_message_dict(self) -> dict[str, Any]:122        """Convert to an OpenAI-format assistant message dict."""123        msg: dict[str, Any] = {"role": "assistant", "content": self.content}124        if self.tool_calls:125            msg["tool_calls"] = [126                {127                    "id": tc.id,128                    "type": "function",129                    "function": {130                        "name": tc.name,131                        "arguments": json.dumps(tc.args),132                    },133                }134                for tc in self.tool_calls135            ]136        return msg137 138 139class LLMClient(ABC):140    """Abstract base for LLM endpoint clients.141 142    Subclass and implement `complete()` for your protocol.143 144    Args:145        endpoint (`str`):146            The `http(s)` base URL of the LLM service (e.g. "http://localhost").147            May include a port and a path (e.g. "http://localhost:8000/v1").148            Credentials, query strings and fragments are rejected.149        port (`int` or `None`):150            The port the service listens on. Appended to `endpoint` when the151            URL does not name one; must match the URL's port when both are given.152    """153 154    def __init__(self, endpoint: str, port: int | None):155        self.endpoint = endpoint156        self.port = port157 158    @abstractmethod159    async def complete(self, prompt: str, **kwargs) -> str:160        """Send a prompt, return the text response.161 162        Args:163            prompt (`str`):164                The user prompt to send.165            **kwargs:166                Override default parameters (temperature, max_tokens, etc.).167 168        Returns:169            The model's text response.170        """171        ...172 173    async def complete_with_tools(174        self,175        messages: list[dict[str, Any]],176        tools: list[dict[str, Any]],177        **kwargs: Any,178    ) -> LLMResponse:179        """Send messages with tool definitions, return a normalized response.180 181        Messages use OpenAI-format dicts (`{"role": "...", "content": "..."}`).182        Tools use MCP tool definitions; they are converted internally.183 184        Args:185            messages (`list[dict[str, Any]]`):186                Conversation history as OpenAI-format message dicts.187            tools (`list[dict[str, Any]]`):188                MCP tool definitions.189            **kwargs:190                Override default parameters (temperature, max_tokens, etc.).191 192        Returns:193            An [`LLMResponse`] with the model's text and any tool calls.194        """195        raise NotImplementedError(196            f"{type(self).__name__} does not support tool calling"197        )198 199    @property200    def base_url(self) -> str:201        """Base URL of the service: `endpoint` plus `port` when the URL names none."""202        return _join_endpoint_port(self.endpoint, self.port)203 204 205class OpenAIClient(LLMClient):206    """Client for OpenAI-compatible APIs.207 208    Works with: OpenAI, vLLM, TGI, Ollama, HuggingFace Inference API,209    or any endpoint that speaks the OpenAI chat completions format.210 211    Args:212        endpoint (`str`):213            The base URL (e.g. "http://localhost"). May include a port and a214            path (e.g. "http://localhost:8000/v1"). The `/v1` API prefix is215            appended when the URL has no path; a URL with a path is used as-is.216        port (`int` or `None`):217            The port number, appended when `endpoint` does not name one; must218            match the URL's port when both are given.219        model (`str`):220            Model name to pass to the API.221        api_key (`str`, *optional*):222            API key. Defaults to "not-needed" for local endpoints.223        system_prompt (`str`, *optional*):224            System message prepended to every request.225        temperature (`float`, *optional*, defaults to `0.0`):226            Default sampling temperature.227        max_tokens (`int`, *optional*, defaults to `256`):228            Default max tokens in the response.229        use_max_completion_tokens (`bool`, *optional*, defaults to `False`):230            Use max_completion_tokens instead of max_tokens. Required for newer OpenAI models231            (gpt-5-mini, o1, o3). Not supported by self-hosted OpenAI-compatible endpoints.232    """233 234    def __init__(235        self,236        endpoint: str,237        port: int | None,238        model: str,239        api_key: str | None = None,240        system_prompt: str | None = None,241        temperature: float = 0.0,242        max_tokens: int = 256,243        use_max_completion_tokens: bool = False,244    ):245        super().__init__(endpoint, port)246        self.model = model247        self.system_prompt = system_prompt248        self.temperature = temperature249        self.max_tokens = max_tokens250        self._tokens_param = (251            "max_completion_tokens" if use_max_completion_tokens else "max_tokens"252        )253        self._omit_temperature = use_max_completion_tokens254 255        self._client = AsyncOpenAI(256            base_url=_openai_base_url(self.base_url),257            api_key=api_key if api_key is not None else "not-needed",258        )259 260    def _chat_completion_kwargs(261        self, messages: list[dict[str, Any]], **kwargs: Any262    ) -> dict[str, Any]:263        create_kwargs: dict[str, Any] = {264            "model": self.model,265            "messages": messages,266            self._tokens_param: kwargs.get("max_tokens", self.max_tokens),267        }268        if not self._omit_temperature:269            create_kwargs["temperature"] = kwargs.get("temperature", self.temperature)270        return create_kwargs271 272    async def complete(self, prompt: str, **kwargs) -> str:273        """Send a chat completion request.274 275        Args:276            prompt (`str`):277                The user message.278            **kwargs:279                Overrides for temperature, max_tokens.280 281        Returns:282            The assistant's response text.283        """284        messages = []285        if self.system_prompt:286            messages.append({"role": "system", "content": self.system_prompt})287        messages.append({"role": "user", "content": prompt})288 289        call_kwargs = self._chat_completion_kwargs(messages, **kwargs)290        response = await self._client.chat.completions.create(**call_kwargs)291        return response.choices[0].message.content or ""292 293    async def complete_with_tools(294        self,295        messages: list[dict[str, Any]],296        tools: list[dict[str, Any]],297        **kwargs: Any,298    ) -> LLMResponse:299        create_kwargs = self._chat_completion_kwargs(messages, **kwargs)300        openai_tools = _mcp_tools_to_openai(tools)301        if openai_tools:302            create_kwargs["tools"] = openai_tools303 304        response = await self._client.chat.completions.create(**create_kwargs)305        msg = response.choices[0].message306 307        tool_calls = []308        if msg.tool_calls:309            for tc in msg.tool_calls:310                tool_calls.append(311                    ToolCall(312                        id=tc.id,313                        name=tc.function.name,314                        args=json.loads(tc.function.arguments),315                    )316                )317 318        return LLMResponse(content=msg.content or "", tool_calls=tool_calls)319 320 321class AnthropicClient(LLMClient):322    """Client for Anthropic's Messages API.323 324    Requires the `anthropic` package (lazy-imported at construction time).325 326    Args:327        endpoint (`str`):328            The base URL (e.g. `https://api.anthropic.com`). May include a port.329        port (`int` or `None`):330            The port number, appended when `endpoint` does not name one; must331            match the URL's port when both are given.332        model (`str`):333            Model name (e.g. "claude-sonnet-4-20250514").334        api_key (`str`, *optional*):335            Anthropic API key.336        system_prompt (`str`, *optional*):337            System message prepended to every request.338        temperature (`float`, *optional*, defaults to `0.0`):339            Default sampling temperature.340        max_tokens (`int`, *optional*, defaults to `256`):341            Default max tokens in the response.342    """343 344    def __init__(345        self,346        endpoint: str,347        port: int | None,348        model: str,349        api_key: str | None = None,350        system_prompt: str | None = None,351        temperature: float = 0.0,352        max_tokens: int = 256,353    ):354        super().__init__(endpoint, port)355        self.model = model356        self.system_prompt = system_prompt357        self.temperature = temperature358        self.max_tokens = max_tokens359 360        try:361            from anthropic import AsyncAnthropic362        except ImportError as exc:363            raise ImportError(364                "AnthropicClient requires the 'anthropic' package. "365                "Install it with: pip install anthropic"366            ) from exc367 368        self._client = AsyncAnthropic(369            base_url=self.base_url,370            api_key=api_key if api_key is not None else "not-needed",371        )372 373    async def complete(self, prompt: str, **kwargs) -> str:374        create_kwargs: dict[str, Any] = {375            "model": self.model,376            "messages": [{"role": "user", "content": prompt}],377            "temperature": kwargs.get("temperature", self.temperature),378            "max_tokens": kwargs.get("max_tokens", self.max_tokens),379        }380        if self.system_prompt:381            create_kwargs["system"] = self.system_prompt382 383        response = await self._client.messages.create(**create_kwargs)384        return "".join(block.text for block in response.content if block.type == "text")385 386    async def complete_with_tools(387        self,388        messages: list[dict[str, Any]],389        tools: list[dict[str, Any]],390        **kwargs: Any,391    ) -> LLMResponse:392        system, anthropic_msgs = _openai_msgs_to_anthropic(messages)393 394        create_kwargs: dict[str, Any] = {395            "model": self.model,396            "messages": anthropic_msgs,397            "temperature": kwargs.get("temperature", self.temperature),398            "max_tokens": kwargs.get("max_tokens", self.max_tokens),399        }400        system_text = system or self.system_prompt401        if system_text:402            create_kwargs["system"] = system_text403        anthropic_tools = _mcp_tools_to_anthropic(tools)404        if anthropic_tools:405            create_kwargs["tools"] = anthropic_tools406 407        response = await self._client.messages.create(**create_kwargs)408 409        content = ""410        tool_calls = []411        for block in response.content:412            if block.type == "text":413                content += block.text414            elif block.type == "tool_use":415                tool_calls.append(416                    ToolCall(id=block.id, name=block.name, args=block.input)417                )418 419        return LLMResponse(content=content, tool_calls=tool_calls)420 421 422# ---------------------------------------------------------------------------423# Factory424# ---------------------------------------------------------------------------425 426_HOSTED_PROVIDERS: dict[str, tuple[str, int, type[LLMClient]]] = {427    "openai": ("https://api.openai.com", 443, OpenAIClient),428    "anthropic": ("https://api.anthropic.com", 443, AnthropicClient),429}430 431# Models that require max_completion_tokens instead of max_tokens and do not432# accept an explicit temperature parameter. Checked by prefix to cover versioned433# names such as "o1-2024-12-17" or "gpt-5-mini-2026-01-15".434_MAX_COMPLETION_TOKENS_PREFIXES: frozenset[str] = frozenset(435    {"gpt-5-mini", "o1", "o3", "o4-mini"}436)437 438 439def create_llm_client(440    provider: str,441    model: str,442    api_key: str,443    *,444    system_prompt: str | None = None,445    temperature: float = 0.0,446    max_tokens: int = 4096,447) -> LLMClient:448    """Create an LLM client for a hosted provider.449 450    Args:451        provider (`str`):452            Provider name ("openai" or "anthropic").453        model (`str`):454            Model identifier.455        api_key (`str`):456            API key for the provider.457        system_prompt (`str`, *optional*):458            System message prepended to every request.459        temperature (`float`, *optional*, defaults to `0.0`):460            Sampling temperature.461        max_tokens (`int`, *optional*, defaults to `4096`):462            Maximum tokens in the response.463 464    Returns:465        A configured [`LLMClient`] instance.466    """467    key = provider.lower()468    if key not in _HOSTED_PROVIDERS:469        raise ValueError(470            f"Unsupported provider: {provider!r}. "471            f"Supported: {sorted(_HOSTED_PROVIDERS)}"472        )473    endpoint, port, cls = _HOSTED_PROVIDERS[key]474    extra: dict[str, Any] = {}475    if cls is OpenAIClient and any(476        model.startswith(prefix) for prefix in _MAX_COMPLETION_TOKENS_PREFIXES477    ):478        extra["use_max_completion_tokens"] = True479    return cls(480        endpoint,481        port,482        model,483        api_key=api_key,484        system_prompt=system_prompt,485        temperature=temperature,486        max_tokens=max_tokens,487        **extra,488    )489 490 491# ---------------------------------------------------------------------------492# MCP tool-schema helpers493# ---------------------------------------------------------------------------494 495 496def _clean_mcp_schema(schema: dict[str, Any]) -> dict[str, Any]:497    """Normalize an MCP tool `inputSchema` for LLM function-calling APIs."""498    if not isinstance(schema, dict):499        return {"type": "object", "properties": {}, "required": []}500 501    # Shallow copy to avoid mutating the caller's schema dict.502    schema = dict(schema)503 504    if "oneOf" in schema:505        for option in schema["oneOf"]:506            if isinstance(option, dict) and option.get("type") == "object":507                schema = option508                break509        else:510            return {"type": "object", "properties": {}, "required": []}511 512    if "allOf" in schema:513        merged: dict[str, Any] = {"type": "object", "properties": {}, "required": []}514        for sub in schema["allOf"]:515            if isinstance(sub, dict):516                if "properties" in sub:517                    merged["properties"].update(sub["properties"])518                if "required" in sub:519                    merged["required"].extend(sub["required"])520        schema = merged521 522    if "anyOf" in schema:523        for option in schema["anyOf"]:524            if isinstance(option, dict) and option.get("type") == "object":525                schema = option526                break527        else:528            return {"type": "object", "properties": {}, "required": []}529 530    schema.setdefault("type", "object")531    if schema.get("type") == "object" and "properties" not in schema:532        schema["properties"] = {}533    return schema534 535 536def _mcp_tools_to_openai(537    mcp_tools: list[dict[str, Any]],538) -> list[dict[str, Any]]:539    """Convert MCP tool definitions to OpenAI function-calling format."""540    result = []541    for tool in mcp_tools:542        input_schema = tool.get(543            "inputSchema", {"type": "object", "properties": {}, "required": []}544        )545        result.append(546            {547                "type": "function",548                "function": {549                    "name": tool["name"],550                    "description": tool.get("description", ""),551                    "parameters": _clean_mcp_schema(input_schema),552                },553            }554        )555    return result556 557 558def _mcp_tools_to_anthropic(559    mcp_tools: list[dict[str, Any]],560) -> list[dict[str, Any]]:561    """Convert MCP tool definitions to Anthropic tool format."""562    result = []563    for tool in mcp_tools:564        input_schema = tool.get(565            "inputSchema", {"type": "object", "properties": {}, "required": []}566        )567        result.append(568            {569                "name": tool["name"],570                "description": tool.get("description", ""),571                "input_schema": _clean_mcp_schema(input_schema),572            }573        )574    return result575 576 577def _openai_msgs_to_anthropic(578    messages: list[dict[str, Any]],579) -> tuple[str, list[dict[str, Any]]]:580    """Convert OpenAI-format messages to Anthropic format.581 582    Returns `(system_text, anthropic_messages)`. System-role messages are583    extracted and concatenated; tool-result messages are converted to584    Anthropic's `tool_result` content blocks inside user turns.585    """586    system_parts: list[str] = []587    anthropic_msgs: list[dict[str, Any]] = []588 589    for msg in messages:590        role = msg["role"]591 592        if role == "system":593            system_parts.append(msg["content"])594 595        elif role == "user":596            anthropic_msgs.append({"role": "user", "content": msg["content"]})597 598        elif role == "assistant":599            if msg.get("tool_calls"):600                content: list[dict[str, Any]] = []601                if msg.get("content"):602                    content.append({"type": "text", "text": msg["content"]})603                for tc in msg["tool_calls"]:604                    args = tc["function"]["arguments"]605                    if isinstance(args, str):606                        args = json.loads(args)607                    content.append(608                        {609                            "type": "tool_use",610                            "id": tc["id"],611                            "name": tc["function"]["name"],612                            "input": args,613                        }614                    )615                anthropic_msgs.append({"role": "assistant", "content": content})616            else:617                anthropic_msgs.append(618                    {"role": "assistant", "content": msg.get("content", "")}619                )620 621        elif role == "tool":622            tool_result = {623                "type": "tool_result",624                "tool_use_id": msg["tool_call_id"],625                "content": msg["content"],626            }627            # Anthropic requires tool results in user turns; merge if possible.628            if (629                anthropic_msgs630                and anthropic_msgs[-1]["role"] == "user"631                and isinstance(anthropic_msgs[-1]["content"], list)632            ):633                anthropic_msgs[-1]["content"].append(tool_result)634            else:635                anthropic_msgs.append({"role": "user", "content": [tool_result]})636 637    system = "\n\n".join(system_parts)638    return system, anthropic_msgs639