mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
Update anyllm.py
Using AnyLLM instead of acompletion
This commit is contained in:
parent
feeb72c0fc
commit
0063fbdd13
1 changed files with 56 additions and 101 deletions
|
|
@ -16,12 +16,12 @@ from .base import Backend, BackendResponse, StreamEvent
|
|||
logger = logging.getLogger(__name__)
|
||||
|
||||
try:
|
||||
from any_llm import acompletion
|
||||
from any_llm import AnyLLM
|
||||
|
||||
ANYLLM_AVAILABLE = True
|
||||
except ImportError:
|
||||
ANYLLM_AVAILABLE = False
|
||||
acompletion = None # type: ignore
|
||||
AnyLLM = None # type: ignore
|
||||
|
||||
|
||||
class AnyLLMBackend(Backend):
|
||||
|
|
@ -43,6 +43,9 @@ class AnyLLMBackend(Backend):
|
|||
self.api_key = api_key
|
||||
self.api_base = api_base
|
||||
|
||||
# Create the AnyLLM instance once and reuse
|
||||
self.llm = AnyLLM.create(self.provider)
|
||||
|
||||
logger.info(f"any-llm backend initialized (provider={provider})")
|
||||
|
||||
@property
|
||||
|
|
@ -57,7 +60,7 @@ class AnyLLMBackend(Backend):
|
|||
"""any-llm supports any model the provider supports."""
|
||||
return True
|
||||
|
||||
def _convert_messages_for_anyllm(self, messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
def _convert_messages(self, messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
"""Convert Anthropic message format to OpenAI/any-llm format."""
|
||||
converted = []
|
||||
for msg in messages:
|
||||
|
|
@ -121,13 +124,13 @@ class AnyLLMBackend(Backend):
|
|||
|
||||
def _to_anthropic_response(
|
||||
self,
|
||||
anyllm_response: Any,
|
||||
response: Any,
|
||||
original_model: str,
|
||||
) -> dict[str, Any]:
|
||||
"""Convert any-llm/OpenAI response to Anthropic format."""
|
||||
msg_id = f"msg_{uuid.uuid4().hex[:24]}"
|
||||
|
||||
choice = anyllm_response.choices[0]
|
||||
choice = response.choices[0]
|
||||
message = choice.message
|
||||
|
||||
content = []
|
||||
|
|
@ -155,10 +158,10 @@ class AnyLLMBackend(Backend):
|
|||
stop_reason = stop_reason_map.get(finish_reason, "end_turn")
|
||||
|
||||
usage = {"input_tokens": 0, "output_tokens": 0}
|
||||
if hasattr(anyllm_response, "usage") and anyllm_response.usage:
|
||||
if hasattr(response, "usage") and response.usage:
|
||||
usage = {
|
||||
"input_tokens": getattr(anyllm_response.usage, "prompt_tokens", 0) or 0,
|
||||
"output_tokens": getattr(anyllm_response.usage, "completion_tokens", 0) or 0,
|
||||
"input_tokens": getattr(response.usage, "prompt_tokens", 0) or 0,
|
||||
"output_tokens": getattr(response.usage, "completion_tokens", 0) or 0,
|
||||
}
|
||||
|
||||
return {
|
||||
|
|
@ -181,14 +184,20 @@ class AnyLLMBackend(Backend):
|
|||
original_model = body.get("model", "gpt-4o")
|
||||
|
||||
try:
|
||||
messages = self._convert_messages_for_anyllm(body.get("messages", []))
|
||||
messages = self._convert_messages(body.get("messages", []))
|
||||
|
||||
kwargs: dict[str, Any] = {
|
||||
"model": original_model,
|
||||
"provider": self.provider,
|
||||
"messages": messages,
|
||||
"stream": False,
|
||||
}
|
||||
# Add system message if present
|
||||
if "system" in body:
|
||||
system = body["system"]
|
||||
if isinstance(system, str):
|
||||
messages.insert(0, {"role": "system", "content": system})
|
||||
elif isinstance(system, list):
|
||||
system_text = " ".join(
|
||||
s.get("text", "") if isinstance(s, dict) else str(s) for s in system
|
||||
)
|
||||
messages.insert(0, {"role": "system", "content": system_text})
|
||||
|
||||
kwargs: dict[str, Any] = {"model": original_model, "messages": messages}
|
||||
|
||||
if "max_tokens" in body:
|
||||
kwargs["max_tokens"] = body["max_tokens"]
|
||||
|
|
@ -199,24 +208,9 @@ class AnyLLMBackend(Backend):
|
|||
if "stop_sequences" in body:
|
||||
kwargs["stop"] = body["stop_sequences"]
|
||||
|
||||
if "system" in body:
|
||||
system = body["system"]
|
||||
if isinstance(system, str):
|
||||
kwargs["messages"].insert(0, {"role": "system", "content": system})
|
||||
elif isinstance(system, list):
|
||||
system_text = " ".join(
|
||||
s.get("text", "") if isinstance(s, dict) else str(s) for s in system
|
||||
)
|
||||
kwargs["messages"].insert(0, {"role": "system", "content": system_text})
|
||||
|
||||
if self.api_key:
|
||||
kwargs["api_key"] = self.api_key
|
||||
if self.api_base:
|
||||
kwargs["api_base"] = self.api_base
|
||||
|
||||
logger.debug(f"any-llm request: provider={self.provider}, model={original_model}")
|
||||
|
||||
response = await acompletion(**kwargs)
|
||||
response = await self.llm.acompletion(**kwargs)
|
||||
anthropic_response = self._to_anthropic_response(response, original_model)
|
||||
|
||||
return BackendResponse(
|
||||
|
|
@ -227,29 +221,7 @@ class AnyLLMBackend(Backend):
|
|||
|
||||
except Exception as e:
|
||||
logger.error(f"any-llm error: {e}")
|
||||
|
||||
error_type = "api_error"
|
||||
status_code = 500
|
||||
|
||||
error_str = str(e).lower()
|
||||
if "authentication" in error_str or "api_key" in error_str or "api key" in error_str:
|
||||
error_type = "authentication_error"
|
||||
status_code = 401
|
||||
elif "rate" in error_str or "limit" in error_str:
|
||||
error_type = "rate_limit_error"
|
||||
status_code = 429
|
||||
elif "not found" in error_str or "model" in error_str:
|
||||
error_type = "not_found_error"
|
||||
status_code = 404
|
||||
|
||||
return BackendResponse(
|
||||
body={
|
||||
"type": "error",
|
||||
"error": {"type": error_type, "message": str(e)},
|
||||
},
|
||||
status_code=status_code,
|
||||
error=str(e),
|
||||
)
|
||||
return self._error_response(e)
|
||||
|
||||
async def stream_message(
|
||||
self,
|
||||
|
|
@ -260,11 +232,15 @@ class AnyLLMBackend(Backend):
|
|||
original_model = body.get("model", "gpt-4o")
|
||||
|
||||
try:
|
||||
messages = self._convert_messages_for_anyllm(body.get("messages", []))
|
||||
messages = self._convert_messages(body.get("messages", []))
|
||||
|
||||
if "system" in body:
|
||||
system = body["system"]
|
||||
if isinstance(system, str):
|
||||
messages.insert(0, {"role": "system", "content": system})
|
||||
|
||||
kwargs: dict[str, Any] = {
|
||||
"model": original_model,
|
||||
"provider": self.provider,
|
||||
"messages": messages,
|
||||
"stream": True,
|
||||
}
|
||||
|
|
@ -273,15 +249,6 @@ class AnyLLMBackend(Backend):
|
|||
kwargs["max_tokens"] = body["max_tokens"]
|
||||
if "temperature" in body:
|
||||
kwargs["temperature"] = body["temperature"]
|
||||
if "system" in body:
|
||||
system = body["system"]
|
||||
if isinstance(system, str):
|
||||
kwargs["messages"].insert(0, {"role": "system", "content": system})
|
||||
|
||||
if self.api_key:
|
||||
kwargs["api_key"] = self.api_key
|
||||
if self.api_base:
|
||||
kwargs["api_base"] = self.api_base
|
||||
|
||||
msg_id = f"msg_{uuid.uuid4().hex[:24]}"
|
||||
|
||||
|
|
@ -311,7 +278,7 @@ class AnyLLMBackend(Backend):
|
|||
},
|
||||
)
|
||||
|
||||
response = await acompletion(**kwargs)
|
||||
response = await self.llm.acompletion(**kwargs)
|
||||
output_tokens = 0
|
||||
|
||||
async for chunk in response:
|
||||
|
|
@ -366,12 +333,7 @@ class AnyLLMBackend(Backend):
|
|||
original_model = body.get("model", "gpt-4o")
|
||||
|
||||
try:
|
||||
kwargs: dict[str, Any] = {
|
||||
"model": original_model,
|
||||
"provider": self.provider,
|
||||
"messages": body.get("messages", []),
|
||||
"stream": False,
|
||||
}
|
||||
kwargs: dict[str, Any] = {"model": original_model, "messages": body.get("messages", [])}
|
||||
|
||||
for param in [
|
||||
"max_tokens",
|
||||
|
|
@ -387,14 +349,9 @@ class AnyLLMBackend(Backend):
|
|||
if param in body:
|
||||
kwargs[param] = body[param]
|
||||
|
||||
if self.api_key:
|
||||
kwargs["api_key"] = self.api_key
|
||||
if self.api_base:
|
||||
kwargs["api_base"] = self.api_base
|
||||
|
||||
logger.debug(f"any-llm OpenAI request: provider={self.provider}, model={original_model}")
|
||||
|
||||
response = await acompletion(**kwargs)
|
||||
response = await self.llm.acompletion(**kwargs)
|
||||
|
||||
response_dict = {
|
||||
"id": response.id,
|
||||
|
|
@ -444,32 +401,30 @@ class AnyLLMBackend(Backend):
|
|||
|
||||
except Exception as e:
|
||||
logger.error(f"any-llm OpenAI error: {e}")
|
||||
return self._error_response(e, openai_format=True)
|
||||
|
||||
error_type = "api_error"
|
||||
status_code = 500
|
||||
def _error_response(self, e: Exception, openai_format: bool = False) -> BackendResponse:
|
||||
"""Build error response from exception."""
|
||||
error_type = "api_error"
|
||||
status_code = 500
|
||||
|
||||
error_str = str(e).lower()
|
||||
if "authentication" in error_str or "api_key" in error_str:
|
||||
error_type = "invalid_api_key"
|
||||
status_code = 401
|
||||
elif "rate" in error_str or "limit" in error_str:
|
||||
error_type = "rate_limit_exceeded"
|
||||
status_code = 429
|
||||
elif "not found" in error_str:
|
||||
error_type = "model_not_found"
|
||||
status_code = 404
|
||||
error_str = str(e).lower()
|
||||
if "authentication" in error_str or "api_key" in error_str or "api key" in error_str:
|
||||
error_type = "invalid_api_key" if openai_format else "authentication_error"
|
||||
status_code = 401
|
||||
elif "rate" in error_str or "limit" in error_str:
|
||||
error_type = "rate_limit_exceeded" if openai_format else "rate_limit_error"
|
||||
status_code = 429
|
||||
elif "not found" in error_str or "model" in error_str:
|
||||
error_type = "model_not_found" if openai_format else "not_found_error"
|
||||
status_code = 404
|
||||
|
||||
return BackendResponse(
|
||||
body={
|
||||
"error": {
|
||||
"message": str(e),
|
||||
"type": error_type,
|
||||
"code": error_type,
|
||||
}
|
||||
},
|
||||
status_code=status_code,
|
||||
error=str(e),
|
||||
)
|
||||
if openai_format:
|
||||
body = {"error": {"message": str(e), "type": error_type, "code": error_type}}
|
||||
else:
|
||||
body = {"type": "error", "error": {"type": error_type, "message": str(e)}}
|
||||
|
||||
return BackendResponse(body=body, status_code=status_code, error=str(e))
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Clean up (no-op for any-llm)."""
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue