Team Ai
Apppublic

openenv/repl

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
llm_client.py507 linesDownload Raw Back to core
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the BSD-style license found in the5# LICENSE file in the root directory of this source tree.6 7"""LLM client abstraction for calling LLM endpoints.8 9Provides a generic RPC abstraction: point it at an endpoint/port, tell it the10protocol, and it works. OpenAI-compatible API is the first implementation,11covering OpenAI, vLLM, TGI, Ollama, HuggingFace Inference API, etc.12Anthropic's native API is supported via ``AnthropicClient``.13 14Usage:15    client = OpenAIClient("http://localhost", 8000, model="meta-llama/...")16    response = await client.complete("What is 2+2?")17 18    # Or use the factory for hosted APIs:19    client = create_llm_client("openai", model="gpt-4", api_key="sk-...")20    response = await client.complete_with_tools(messages, tools)21"""22 23from __future__ import annotations24 25import json26from abc import ABC, abstractmethod27from dataclasses import dataclass, field28from typing import Any29 30from openai import AsyncOpenAI31 32 33@dataclass34class ToolCall:35    """A single tool/function call returned by the model."""36 37    id: str38    name: str39    args: dict[str, Any]40 41 42@dataclass43class LLMResponse:44    """Normalized response from an LLM, with optional tool calls."""45 46    content: str47    tool_calls: list[ToolCall] = field(default_factory=list)48 49    def to_message_dict(self) -> dict[str, Any]:50        """Convert to an OpenAI-format assistant message dict."""51        msg: dict[str, Any] = {"role": "assistant", "content": self.content}52        if self.tool_calls:53            msg["tool_calls"] = [54                {55                    "id": tc.id,56                    "type": "function",57                    "function": {58                        "name": tc.name,59                        "arguments": json.dumps(tc.args),60                    },61                }62                for tc in self.tool_calls63            ]64        return msg65 66 67class LLMClient(ABC):68    """Abstract base for LLM endpoint clients.69 70    Subclass and implement ``complete()`` for your protocol.71 72    Args:73        endpoint: The base URL of the LLM service (e.g. "http://localhost").74        port: The port the service listens on.75    """76 77    def __init__(self, endpoint: str, port: int):78        self.endpoint = endpoint79        self.port = port80 81    @abstractmethod82    async def complete(self, prompt: str, **kwargs) -> str:83        """Send a prompt, return the text response.84 85        Args:86            prompt: The user prompt to send.87            **kwargs: Override default parameters (temperature, max_tokens, etc.).88 89        Returns:90            The model's text response.91        """92        ...93 94    async def complete_with_tools(95        self,96        messages: list[dict[str, Any]],97        tools: list[dict[str, Any]],98        **kwargs: Any,99    ) -> LLMResponse:100        """Send messages with tool definitions, return a normalized response.101 102        Messages use OpenAI-format dicts (``{"role": "...", "content": "..."}``).103        Tools use MCP tool definitions; they are converted internally.104 105        Args:106            messages: Conversation history as OpenAI-format message dicts.107            tools: MCP tool definitions.108            **kwargs: Override default parameters (temperature, max_tokens, etc.).109 110        Returns:111            An ``LLMResponse`` with the model's text and any tool calls.112        """113        raise NotImplementedError(114            f"{type(self).__name__} does not support tool calling"115        )116 117    @property118    def base_url(self) -> str:119        """Construct base URL from endpoint and port."""120        return f"{self.endpoint}:{self.port}"121 122 123class OpenAIClient(LLMClient):124    """Client for OpenAI-compatible APIs.125 126    Works with: OpenAI, vLLM, TGI, Ollama, HuggingFace Inference API,127    or any endpoint that speaks the OpenAI chat completions format.128 129    Args:130        endpoint: The base URL (e.g. "http://localhost").131        port: The port number.132        model: Model name to pass to the API.133        api_key: API key. Defaults to "not-needed" for local endpoints.134        system_prompt: Optional system message prepended to every request.135        temperature: Default sampling temperature.136        max_tokens: Default max tokens in the response.137    """138 139    def __init__(140        self,141        endpoint: str,142        port: int,143        model: str,144        api_key: str | None = None,145        system_prompt: str | None = None,146        temperature: float = 0.0,147        max_tokens: int = 256,148    ):149        super().__init__(endpoint, port)150        self.model = model151        self.system_prompt = system_prompt152        self.temperature = temperature153        self.max_tokens = max_tokens154 155        self._client = AsyncOpenAI(156            base_url=f"{self.base_url}/v1",157            api_key=api_key if api_key is not None else "not-needed",158        )159 160    async def complete(self, prompt: str, **kwargs) -> str:161        """Send a chat completion request.162 163        Args:164            prompt: The user message.165            **kwargs: Overrides for temperature, max_tokens.166 167        Returns:168            The assistant's response text.169        """170        messages = []171        if self.system_prompt:172            messages.append({"role": "system", "content": self.system_prompt})173        messages.append({"role": "user", "content": prompt})174 175        response = await self._client.chat.completions.create(176            model=self.model,177            messages=messages,178            temperature=kwargs.get("temperature", self.temperature),179            max_tokens=kwargs.get("max_tokens", self.max_tokens),180        )181        return response.choices[0].message.content or ""182 183    async def complete_with_tools(184        self,185        messages: list[dict[str, Any]],186        tools: list[dict[str, Any]],187        **kwargs: Any,188    ) -> LLMResponse:189        create_kwargs: dict[str, Any] = {190            "model": self.model,191            "messages": messages,192            "temperature": kwargs.get("temperature", self.temperature),193            "max_tokens": kwargs.get("max_tokens", self.max_tokens),194        }195        openai_tools = _mcp_tools_to_openai(tools)196        if openai_tools:197            create_kwargs["tools"] = openai_tools198 199        response = await self._client.chat.completions.create(**create_kwargs)200        msg = response.choices[0].message201 202        tool_calls = []203        if msg.tool_calls:204            for tc in msg.tool_calls:205                tool_calls.append(206                    ToolCall(207                        id=tc.id,208                        name=tc.function.name,209                        args=json.loads(tc.function.arguments),210                    )211                )212 213        return LLMResponse(content=msg.content or "", tool_calls=tool_calls)214 215 216class AnthropicClient(LLMClient):217    """Client for Anthropic's Messages API.218 219    Requires the ``anthropic`` package (lazy-imported at construction time).220 221    Args:222        endpoint: The base URL (e.g. "https://api.anthropic.com").223        port: The port number.224        model: Model name (e.g. "claude-sonnet-4-20250514").225        api_key: Anthropic API key.226        system_prompt: Optional system message prepended to every request.227        temperature: Default sampling temperature.228        max_tokens: Default max tokens in the response.229    """230 231    def __init__(232        self,233        endpoint: str,234        port: int,235        model: str,236        api_key: str | None = None,237        system_prompt: str | None = None,238        temperature: float = 0.0,239        max_tokens: int = 256,240    ):241        super().__init__(endpoint, port)242        self.model = model243        self.system_prompt = system_prompt244        self.temperature = temperature245        self.max_tokens = max_tokens246 247        try:248            from anthropic import AsyncAnthropic249        except ImportError as exc:250            raise ImportError(251                "AnthropicClient requires the 'anthropic' package. "252                "Install it with: pip install anthropic"253            ) from exc254 255        self._client = AsyncAnthropic(256            base_url=self.base_url,257            api_key=api_key if api_key is not None else "not-needed",258        )259 260    async def complete(self, prompt: str, **kwargs) -> str:261        create_kwargs: dict[str, Any] = {262            "model": self.model,263            "messages": [{"role": "user", "content": prompt}],264            "temperature": kwargs.get("temperature", self.temperature),265            "max_tokens": kwargs.get("max_tokens", self.max_tokens),266        }267        if self.system_prompt:268            create_kwargs["system"] = self.system_prompt269 270        response = await self._client.messages.create(**create_kwargs)271        return "".join(block.text for block in response.content if block.type == "text")272 273    async def complete_with_tools(274        self,275        messages: list[dict[str, Any]],276        tools: list[dict[str, Any]],277        **kwargs: Any,278    ) -> LLMResponse:279        system, anthropic_msgs = _openai_msgs_to_anthropic(messages)280 281        create_kwargs: dict[str, Any] = {282            "model": self.model,283            "messages": anthropic_msgs,284            "temperature": kwargs.get("temperature", self.temperature),285            "max_tokens": kwargs.get("max_tokens", self.max_tokens),286        }287        system_text = system or self.system_prompt288        if system_text:289            create_kwargs["system"] = system_text290        anthropic_tools = _mcp_tools_to_anthropic(tools)291        if anthropic_tools:292            create_kwargs["tools"] = anthropic_tools293 294        response = await self._client.messages.create(**create_kwargs)295 296        content = ""297        tool_calls = []298        for block in response.content:299            if block.type == "text":300                content += block.text301            elif block.type == "tool_use":302                tool_calls.append(303                    ToolCall(id=block.id, name=block.name, args=block.input)304                )305 306        return LLMResponse(content=content, tool_calls=tool_calls)307 308 309# ---------------------------------------------------------------------------310# Factory311# ---------------------------------------------------------------------------312 313_HOSTED_PROVIDERS: dict[str, tuple[str, int, type[LLMClient]]] = {314    "openai": ("https://api.openai.com", 443, OpenAIClient),315    "anthropic": ("https://api.anthropic.com", 443, AnthropicClient),316}317 318 319def create_llm_client(320    provider: str,321    model: str,322    api_key: str,323    *,324    system_prompt: str | None = None,325    temperature: float = 0.0,326    max_tokens: int = 4096,327) -> LLMClient:328    """Create an LLM client for a hosted provider.329 330    Args:331        provider: Provider name ("openai" or "anthropic").332        model: Model identifier.333        api_key: API key for the provider.334        system_prompt: Optional system message prepended to every request.335        temperature: Sampling temperature.336        max_tokens: Maximum tokens in the response.337 338    Returns:339        A configured ``LLMClient`` instance.340    """341    key = provider.lower()342    if key not in _HOSTED_PROVIDERS:343        raise ValueError(344            f"Unsupported provider: {provider!r}. "345            f"Supported: {sorted(_HOSTED_PROVIDERS)}"346        )347    endpoint, port, cls = _HOSTED_PROVIDERS[key]348    return cls(349        endpoint,350        port,351        model,352        api_key=api_key,353        system_prompt=system_prompt,354        temperature=temperature,355        max_tokens=max_tokens,356    )357 358 359# ---------------------------------------------------------------------------360# MCP tool-schema helpers361# ---------------------------------------------------------------------------362 363 364def _clean_mcp_schema(schema: dict[str, Any]) -> dict[str, Any]:365    """Normalize an MCP tool ``inputSchema`` for LLM function-calling APIs."""366    if not isinstance(schema, dict):367        return {"type": "object", "properties": {}, "required": []}368 369    # Shallow copy to avoid mutating the caller's schema dict.370    schema = dict(schema)371 372    if "oneOf" in schema:373        for option in schema["oneOf"]:374            if isinstance(option, dict) and option.get("type") == "object":375                schema = option376                break377        else:378            return {"type": "object", "properties": {}, "required": []}379 380    if "allOf" in schema:381        merged: dict[str, Any] = {"type": "object", "properties": {}, "required": []}382        for sub in schema["allOf"]:383            if isinstance(sub, dict):384                if "properties" in sub:385                    merged["properties"].update(sub["properties"])386                if "required" in sub:387                    merged["required"].extend(sub["required"])388        schema = merged389 390    if "anyOf" in schema:391        for option in schema["anyOf"]:392            if isinstance(option, dict) and option.get("type") == "object":393                schema = option394                break395        else:396            return {"type": "object", "properties": {}, "required": []}397 398    schema.setdefault("type", "object")399    if schema.get("type") == "object" and "properties" not in schema:400        schema["properties"] = {}401    return schema402 403 404def _mcp_tools_to_openai(405    mcp_tools: list[dict[str, Any]],406) -> list[dict[str, Any]]:407    """Convert MCP tool definitions to OpenAI function-calling format."""408    result = []409    for tool in mcp_tools:410        input_schema = tool.get(411            "inputSchema", {"type": "object", "properties": {}, "required": []}412        )413        result.append(414            {415                "type": "function",416                "function": {417                    "name": tool["name"],418                    "description": tool.get("description", ""),419                    "parameters": _clean_mcp_schema(input_schema),420                },421            }422        )423    return result424 425 426def _mcp_tools_to_anthropic(427    mcp_tools: list[dict[str, Any]],428) -> list[dict[str, Any]]:429    """Convert MCP tool definitions to Anthropic tool format."""430    result = []431    for tool in mcp_tools:432        input_schema = tool.get(433            "inputSchema", {"type": "object", "properties": {}, "required": []}434        )435        result.append(436            {437                "name": tool["name"],438                "description": tool.get("description", ""),439                "input_schema": _clean_mcp_schema(input_schema),440            }441        )442    return result443 444 445def _openai_msgs_to_anthropic(446    messages: list[dict[str, Any]],447) -> tuple[str, list[dict[str, Any]]]:448    """Convert OpenAI-format messages to Anthropic format.449 450    Returns ``(system_text, anthropic_messages)``.  System-role messages are451    extracted and concatenated; tool-result messages are converted to452    Anthropic's ``tool_result`` content blocks inside user turns.453    """454    system_parts: list[str] = []455    anthropic_msgs: list[dict[str, Any]] = []456 457    for msg in messages:458        role = msg["role"]459 460        if role == "system":461            system_parts.append(msg["content"])462 463        elif role == "user":464            anthropic_msgs.append({"role": "user", "content": msg["content"]})465 466        elif role == "assistant":467            if msg.get("tool_calls"):468                content: list[dict[str, Any]] = []469                if msg.get("content"):470                    content.append({"type": "text", "text": msg["content"]})471                for tc in msg["tool_calls"]:472                    args = tc["function"]["arguments"]473                    if isinstance(args, str):474                        args = json.loads(args)475                    content.append(476                        {477                            "type": "tool_use",478                            "id": tc["id"],479                            "name": tc["function"]["name"],480                            "input": args,481                        }482                    )483                anthropic_msgs.append({"role": "assistant", "content": content})484            else:485                anthropic_msgs.append(486                    {"role": "assistant", "content": msg.get("content", "")}487                )488 489        elif role == "tool":490            tool_result = {491                "type": "tool_result",492                "tool_use_id": msg["tool_call_id"],493                "content": msg["content"],494            }495            # Anthropic requires tool results in user turns; merge if possible.496            if (497                anthropic_msgs498                and anthropic_msgs[-1]["role"] == "user"499                and isinstance(anthropic_msgs[-1]["content"], list)500            ):501                anthropic_msgs[-1]["content"].append(tool_result)502            else:503                anthropic_msgs.append({"role": "user", "content": [tool_result]})504 505    system = "\n\n".join(system_parts)506    return system, anthropic_msgs507