openenv/atari_env
3
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 