mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-14 04:36:52 +08:00
## What This pull request adds **MWS GPT Model Hub** as a built-in model provider in RAGFlow. The integration allows users to configure an MWS project endpoint and token, discover the models available to that project, and use supported MWS models for chat completion, embeddings, and reranking. Co-authored-by: ilarionov_n <ilarionov_n@promis.ru>
2611 lines
111 KiB
Python
2611 lines
111 KiB
Python
#
|
|
# Copyright 2025 The InfiniFlow Authors. All Rights Reserved.
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
#
|
|
import asyncio
|
|
import json
|
|
import logging
|
|
import os
|
|
import random
|
|
import re
|
|
import time
|
|
from abc import ABC
|
|
from copy import deepcopy
|
|
from urllib.parse import urljoin
|
|
|
|
import aiohttp
|
|
import json_repair
|
|
from json.decoder import JSONDecodeError
|
|
import litellm
|
|
from openai import AsyncOpenAI, OpenAI
|
|
from enum import StrEnum
|
|
|
|
from common.aimlapi_utils import attribution_headers
|
|
from common.misc_utils import thread_pool_exec
|
|
from common.llm_request_context import current_llm_user
|
|
from common.token_utils import num_tokens_from_string, total_token_count_from_response, usage_from_response
|
|
from rag.llm import FACTORY_DEFAULT_BASE_URL, LITELLM_PROVIDER_PREFIX, SupportedLiteLLMProvider
|
|
from rag.llm.key_utils import _normalize_replicate_key
|
|
from rag.llm.mws_utils import mws_api_url, require_mws_token
|
|
from rag.llm.tool_decorator import FunctionToolSession, is_tool
|
|
from rag.nlp import is_chinese, is_english
|
|
from rag.utils.url_utils import ensure_v1
|
|
|
|
|
|
class LLMErrorCode(StrEnum):
|
|
ERROR_RATE_LIMIT = "RATE_LIMIT_EXCEEDED"
|
|
ERROR_AUTHENTICATION = "AUTH_ERROR"
|
|
ERROR_INVALID_REQUEST = "INVALID_REQUEST"
|
|
ERROR_SERVER = "SERVER_ERROR"
|
|
ERROR_TIMEOUT = "TIMEOUT"
|
|
ERROR_CONNECTION = "CONNECTION_ERROR"
|
|
ERROR_MODEL = "MODEL_ERROR"
|
|
ERROR_MAX_ROUNDS = "ERROR_MAX_ROUNDS"
|
|
ERROR_CONTENT_FILTER = "CONTENT_FILTERED"
|
|
ERROR_QUOTA = "QUOTA_EXCEEDED"
|
|
ERROR_MAX_RETRIES = "MAX_RETRIES_EXCEEDED"
|
|
ERROR_GENERIC = "GENERIC_ERROR"
|
|
|
|
|
|
class ReActMode(StrEnum):
|
|
FUNCTION_CALL = "function_call"
|
|
REACT = "react"
|
|
|
|
|
|
ERROR_PREFIX = "**ERROR**"
|
|
LENGTH_NOTIFICATION_CN = "······\n由于大模型的上下文窗口大小限制,回答已经被大模型截断。"
|
|
LENGTH_NOTIFICATION_EN = "...\nThe answer is truncated by your chosen LLM due to its limitation on context length."
|
|
|
|
# Generation parameters that are safe to forward to the underlying completion
|
|
# call. `gen_conf` originates from a chat assistant's `llm_setting`, which can
|
|
# also carry RAGFlow-internal metadata (e.g. `model_type`). Anything outside
|
|
# this set is dropped so providers don't reject the request with errors like
|
|
# "Extra inputs are not permitted" / "Unknown parameter: 'model_type'" (#15427).
|
|
ALLOWED_GEN_CONF_KEYS = frozenset(
|
|
{
|
|
"temperature",
|
|
"max_completion_tokens",
|
|
"top_p",
|
|
"stream",
|
|
"stream_options",
|
|
"stop",
|
|
"n",
|
|
"presence_penalty",
|
|
"frequency_penalty",
|
|
"functions",
|
|
"function_call",
|
|
"logit_bias",
|
|
"user",
|
|
"response_format",
|
|
"seed",
|
|
"tools",
|
|
"tool_choice",
|
|
"logprobs",
|
|
"top_logprobs",
|
|
"extra_headers",
|
|
}
|
|
)
|
|
|
|
# LiteLLM additionally understands reasoning-control parameters that the
|
|
# model-family policies may inject into `gen_conf` (e.g. `thinking` for
|
|
# Anthropic / Kimi reasoning models, `enable_thinking` for Qwen models,
|
|
# `reasoning_effort` for OpenAI o-series).
|
|
LITELLM_ALLOWED_GEN_CONF_KEYS = ALLOWED_GEN_CONF_KEYS | frozenset(
|
|
{
|
|
"thinking",
|
|
"enable_thinking",
|
|
"reasoning_effort",
|
|
"extra_body",
|
|
}
|
|
)
|
|
|
|
|
|
def _apply_model_family_policies(
|
|
model_name: str,
|
|
*,
|
|
backend: str,
|
|
provider: SupportedLiteLLMProvider | str | None = None,
|
|
gen_conf: dict | None = None,
|
|
request_kwargs: dict | None = None,
|
|
):
|
|
model_name_lower = (model_name or "").lower()
|
|
sanitized_gen_conf = deepcopy(gen_conf) if gen_conf else {}
|
|
sanitized_kwargs = dict(request_kwargs) if request_kwargs else {}
|
|
|
|
def _thinking_type():
|
|
val = sanitized_gen_conf.get("thinking")
|
|
if isinstance(val, dict):
|
|
val = val.get("type")
|
|
|
|
enable_thinking = sanitized_gen_conf.get("enable_thinking")
|
|
|
|
if isinstance(val, str) and val in {"enabled", "disabled"}:
|
|
return val
|
|
if isinstance(enable_thinking, bool):
|
|
return "enabled" if enable_thinking else "disabled"
|
|
return None
|
|
|
|
def _pop_thinking_controls():
|
|
sanitized_gen_conf.pop("thinking", None)
|
|
sanitized_gen_conf.pop("enable_thinking", None)
|
|
|
|
def _merge_extra_body(target: dict, extra: dict) -> None:
|
|
body = target.get("extra_body")
|
|
if not isinstance(body, dict):
|
|
body = {}
|
|
body.update(extra)
|
|
target["extra_body"] = body
|
|
|
|
thinking_type = _thinking_type()
|
|
|
|
# Qwen3 keeps RAGFlow's system default of disabling thinking unless explicitly overridden.
|
|
if "qwen3" in model_name_lower:
|
|
_pop_thinking_controls()
|
|
# -preview variants (e.g. qwen3.8-max-preview) only accept
|
|
# enable_thinking=True; the API rejects any other value.
|
|
if "-preview" in model_name_lower:
|
|
enable_thinking = True
|
|
else:
|
|
enable_thinking = thinking_type == "enabled" if thinking_type else False
|
|
if backend == "litellm" and provider in {
|
|
SupportedLiteLLMProvider.Tongyi_Qianwen,
|
|
SupportedLiteLLMProvider.Dashscope,
|
|
}:
|
|
sanitized_gen_conf["enable_thinking"] = enable_thinking
|
|
else:
|
|
_merge_extra_body(sanitized_kwargs, {"enable_thinking": enable_thinking})
|
|
|
|
if backend == "base":
|
|
return sanitized_gen_conf, sanitized_kwargs
|
|
|
|
if backend == "litellm":
|
|
if provider in {SupportedLiteLLMProvider.OpenAI, SupportedLiteLLMProvider.Azure_OpenAI} and "gpt-5" in model_name_lower:
|
|
for key in ("temperature", "top_p", "logprobs", "top_logprobs"):
|
|
sanitized_gen_conf.pop(key, None)
|
|
sanitized_kwargs.pop(key, None)
|
|
elif provider == SupportedLiteLLMProvider.Anthropic and model_name_lower in {"claude-opus-4-7", "claude-opus-4-8"}:
|
|
for key in ("temperature", "top_p", "top_k"):
|
|
sanitized_gen_conf.pop(key, None)
|
|
sanitized_kwargs.pop(key, None)
|
|
|
|
if provider == SupportedLiteLLMProvider.HunYuan:
|
|
for key in ("presence_penalty", "frequency_penalty"):
|
|
sanitized_gen_conf.pop(key, None)
|
|
elif provider == SupportedLiteLLMProvider.Moonshot:
|
|
if thinking_type:
|
|
_pop_thinking_controls()
|
|
sanitized_gen_conf["thinking"] = {"type": thinking_type}
|
|
|
|
if thinking_type or "kimi-k2.5" in model_name_lower or "kimi-k2.6" in model_name_lower:
|
|
sanitized_gen_conf.pop("temperature", None)
|
|
sanitized_gen_conf["top_p"] = 0.95
|
|
sanitized_gen_conf["n"] = 1
|
|
sanitized_gen_conf["presence_penalty"] = 0.0
|
|
sanitized_gen_conf["frequency_penalty"] = 0.0
|
|
elif provider == SupportedLiteLLMProvider.ZHIPU_AI and "glm" in model_name_lower and thinking_type:
|
|
_pop_thinking_controls()
|
|
sanitized_gen_conf["thinking"] = {"type": thinking_type}
|
|
|
|
return sanitized_gen_conf, sanitized_kwargs
|
|
|
|
return sanitized_gen_conf, sanitized_kwargs
|
|
|
|
|
|
def _move_litellm_provider_body_fields(provider: SupportedLiteLLMProvider | str | None, completion_args: dict) -> dict:
|
|
provider_body_fields = {
|
|
SupportedLiteLLMProvider.Tongyi_Qianwen: {"enable_thinking"},
|
|
SupportedLiteLLMProvider.Dashscope: {"enable_thinking"},
|
|
SupportedLiteLLMProvider.Moonshot: {"thinking"},
|
|
SupportedLiteLLMProvider.ZHIPU_AI: {"thinking"},
|
|
}.get(provider, set())
|
|
|
|
body = completion_args.get("extra_body")
|
|
if not isinstance(body, dict):
|
|
body = {}
|
|
moved = False
|
|
for key in provider_body_fields:
|
|
if key in completion_args:
|
|
body[key] = completion_args.pop(key)
|
|
moved = True
|
|
if moved or body:
|
|
completion_args["extra_body"] = body
|
|
return completion_args
|
|
|
|
|
|
class Base(ABC):
|
|
def __init__(self, key, model_name, base_url, **kwargs):
|
|
timeout = int(os.environ.get("LLM_TIMEOUT_SECONDS", 600))
|
|
self.base_url = ensure_v1(base_url)
|
|
self.client = OpenAI(api_key=key, base_url=self.base_url, timeout=timeout)
|
|
self.async_client = AsyncOpenAI(api_key=key, base_url=self.base_url, timeout=timeout)
|
|
self.model_name = model_name
|
|
# Configure retry parameters
|
|
self.max_retries = kwargs.get("max_retries", int(os.environ.get("LLM_MAX_RETRIES", 5)))
|
|
self.base_delay = kwargs.get("retry_interval", float(os.environ.get("LLM_BASE_DELAY", 2.0)))
|
|
self.max_rounds = kwargs.get("max_rounds", 5)
|
|
self.is_tools = False
|
|
self.tools = []
|
|
self.toolcall_sessions = {}
|
|
# Token usage split (prompt/completion/total) of the most recent chat call.
|
|
# Consumed by LLMBundle for accurate Langfuse reporting and run aggregation.
|
|
self.last_usage = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
|
|
|
def _get_delay(self):
|
|
return self.base_delay * random.uniform(10, 150)
|
|
|
|
def _classify_error(self, error):
|
|
error_str = str(error).lower()
|
|
|
|
keywords_mapping = [
|
|
(["quota", "capacity", "credit", "billing", "balance", "欠费"], LLMErrorCode.ERROR_QUOTA),
|
|
(["rate limit", "429", "tpm limit", "too many requests", "requests per minute"], LLMErrorCode.ERROR_RATE_LIMIT),
|
|
(["auth", "key", "apikey", "401", "forbidden", "permission"], LLMErrorCode.ERROR_AUTHENTICATION),
|
|
(["invalid", "bad request", "400", "format", "malformed", "parameter"], LLMErrorCode.ERROR_INVALID_REQUEST),
|
|
(["server", "503", "502", "504", "500", "unavailable"], LLMErrorCode.ERROR_SERVER),
|
|
(["timeout", "timed out"], LLMErrorCode.ERROR_TIMEOUT),
|
|
(["connect", "network", "unreachable", "dns"], LLMErrorCode.ERROR_CONNECTION),
|
|
(["filter", "content", "policy", "blocked", "safety", "inappropriate"], LLMErrorCode.ERROR_CONTENT_FILTER),
|
|
(["model", "not found", "does not exist", "not available"], LLMErrorCode.ERROR_MODEL),
|
|
(["max rounds"], LLMErrorCode.ERROR_MODEL),
|
|
]
|
|
for words, code in keywords_mapping:
|
|
if re.search("({})".format("|".join(words)), error_str):
|
|
return code
|
|
|
|
return LLMErrorCode.ERROR_GENERIC
|
|
|
|
def _clean_conf(self, gen_conf):
|
|
if "max_tokens" in gen_conf:
|
|
del gen_conf["max_tokens"]
|
|
|
|
gen_conf = {k: v for k, v in gen_conf.items() if k in ALLOWED_GEN_CONF_KEYS}
|
|
return gen_conf
|
|
|
|
async def _async_chat_streamly(self, history, gen_conf, **kwargs):
|
|
logging.info("[HISTORY STREAMLY]" + json.dumps(history, ensure_ascii=False, indent=4))
|
|
reasoning_start = False
|
|
|
|
gen_conf, extra_request_kwargs = _apply_model_family_policies(
|
|
self.model_name,
|
|
backend="base",
|
|
gen_conf=gen_conf,
|
|
request_kwargs={},
|
|
)
|
|
request_kwargs = {"model": self.model_name, "messages": history, "stream": True, **gen_conf}
|
|
stop = kwargs.get("stop")
|
|
if stop:
|
|
request_kwargs["stop"] = stop
|
|
request_kwargs.update(extra_request_kwargs)
|
|
|
|
response = await self.async_client.chat.completions.create(**request_kwargs)
|
|
async for resp in response:
|
|
if not resp.choices:
|
|
continue
|
|
if not resp.choices[0].delta.content:
|
|
resp.choices[0].delta.content = ""
|
|
_reasoning = getattr(resp.choices[0].delta, "reasoning_content", None) or getattr(resp.choices[0].delta, "reasoning", None)
|
|
if kwargs.get("with_reasoning", True) and _reasoning:
|
|
ans = ""
|
|
if not reasoning_start:
|
|
reasoning_start = True
|
|
ans = "<think>"
|
|
ans += _reasoning + "</think>"
|
|
else:
|
|
reasoning_start = False
|
|
ans = resp.choices[0].delta.content
|
|
tol = total_token_count_from_response(resp)
|
|
if not tol:
|
|
tol = num_tokens_from_string(resp.choices[0].delta.content)
|
|
|
|
finish_reason = resp.choices[0].finish_reason if hasattr(resp.choices[0], "finish_reason") else ""
|
|
if finish_reason == "length":
|
|
if is_chinese(ans):
|
|
ans += LENGTH_NOTIFICATION_CN
|
|
else:
|
|
ans += LENGTH_NOTIFICATION_EN
|
|
yield ans, tol
|
|
|
|
async def async_chat_streamly(self, system, history, gen_conf: dict | None = None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
if system and history and history[0].get("role") != "system":
|
|
history.insert(0, {"role": "system", "content": system})
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
ans = ""
|
|
total_tokens = 0
|
|
# Reset so a stale split from a previous call can't leak into this one.
|
|
self.last_usage = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
|
|
|
for attempt in range(self.max_retries + 1):
|
|
try:
|
|
async for delta_ans, tol in self._async_chat_streamly(history, gen_conf, **kwargs):
|
|
ans = delta_ans
|
|
total_tokens += tol
|
|
yield ans
|
|
|
|
yield total_tokens
|
|
return
|
|
except Exception as e:
|
|
e = await self._exceptions_async(e, attempt)
|
|
if e:
|
|
yield e
|
|
yield total_tokens
|
|
return
|
|
|
|
def _length_stop(self, ans):
|
|
if is_chinese([ans]):
|
|
return ans + LENGTH_NOTIFICATION_CN
|
|
return ans + LENGTH_NOTIFICATION_EN
|
|
|
|
@property
|
|
def _retryable_errors(self) -> set[str]:
|
|
return {
|
|
LLMErrorCode.ERROR_RATE_LIMIT,
|
|
LLMErrorCode.ERROR_SERVER,
|
|
}
|
|
|
|
def _should_retry(self, error_code: str) -> bool:
|
|
return error_code in self._retryable_errors
|
|
|
|
def _exceptions(self, e, attempt) -> str | None:
|
|
logging.exception("OpenAI chat_with_tools")
|
|
# Classify the error
|
|
error_code = self._classify_error(e)
|
|
if attempt == self.max_retries:
|
|
error_code = LLMErrorCode.ERROR_MAX_RETRIES
|
|
|
|
if self._should_retry(error_code):
|
|
delay = self._get_delay()
|
|
logging.warning(f"Error: {error_code}. Retrying in {delay:.2f} seconds... (Attempt {attempt + 1}/{self.max_retries})")
|
|
time.sleep(delay)
|
|
return None
|
|
|
|
msg = f"{ERROR_PREFIX}: {error_code} - {str(e)}"
|
|
logging.error(f"sync base giving up: {msg}")
|
|
return msg
|
|
|
|
async def _exceptions_async(self, e, attempt):
|
|
logging.exception("OpenAI async completion")
|
|
error_code = self._classify_error(e)
|
|
if attempt == self.max_retries:
|
|
error_code = LLMErrorCode.ERROR_MAX_RETRIES
|
|
|
|
if self._should_retry(error_code):
|
|
delay = self._get_delay()
|
|
logging.warning(f"Error: {error_code}. Retrying in {delay:.2f} seconds... (Attempt {attempt + 1}/{self.max_retries})")
|
|
await asyncio.sleep(delay)
|
|
return None
|
|
|
|
msg = f"{ERROR_PREFIX}: {error_code} - {str(e)}"
|
|
logging.error(f"async base giving up: {msg}")
|
|
return msg
|
|
|
|
def _verbose_tool_use(self, name, args, res):
|
|
return "<tool_call>" + json.dumps({"name": name, "args": args, "result": res}, ensure_ascii=False, indent=2) + "</tool_call>"
|
|
|
|
def _append_history(self, hist, tool_call, tool_res):
|
|
hist.append(
|
|
{
|
|
"role": "assistant",
|
|
"content": None,
|
|
"tool_calls": [
|
|
{
|
|
"index": getattr(tool_call, "index", None),
|
|
"id": tool_call.id,
|
|
"function": {
|
|
"name": tool_call.function.name,
|
|
"arguments": tool_call.function.arguments,
|
|
},
|
|
"type": "function",
|
|
},
|
|
],
|
|
}
|
|
)
|
|
try:
|
|
if isinstance(tool_res, dict):
|
|
tool_res = json.dumps(tool_res, ensure_ascii=False)
|
|
finally:
|
|
hist.append({"role": "tool", "tool_call_id": tool_call.id, "content": str(tool_res)})
|
|
return hist
|
|
|
|
def _append_history_batch(self, hist, results):
|
|
"""
|
|
Append a batch of tool calls to history following the OpenAI protocol:
|
|
one assistant message containing all tool_calls, followed by one tool message per call.
|
|
results: list of (tool_call, name, args, result, error)
|
|
"""
|
|
hist.append(
|
|
{
|
|
"role": "assistant",
|
|
"content": None,
|
|
"tool_calls": [
|
|
{
|
|
"index": getattr(tc, "index", None),
|
|
"id": tc.id,
|
|
"function": {"name": tc.function.name, "arguments": tc.function.arguments},
|
|
"type": "function",
|
|
}
|
|
for tc, _, _, _, _ in results
|
|
],
|
|
}
|
|
)
|
|
for tc, _, _, result, err in results:
|
|
if err:
|
|
content = str(err)
|
|
elif isinstance(result, dict):
|
|
content = json.dumps(result, ensure_ascii=False)
|
|
else:
|
|
content = str(result)
|
|
hist.append({"role": "tool", "tool_call_id": tc.id, "content": content})
|
|
return hist
|
|
|
|
def bind_tools(self, toolcall_session=None, tools=None):
|
|
"""Register tools the LLM can call.
|
|
|
|
Two calling styles are accepted:
|
|
|
|
* Legacy: ``bind_tools(toolcall_session, tools_schemas)`` where
|
|
``toolcall_session`` implements :class:`ToolCallSession` and
|
|
``tools_schemas`` is a pre-built list of OpenAI function-schema
|
|
dicts (used by the agent/dialog layer).
|
|
* Decorator: ``bind_tools(tools=[fn1, fn2, ...])`` where each ``fn``
|
|
is decorated with :func:`rag.llm.tool_decorator.tool`. The session
|
|
and schemas are derived from the callables automatically.
|
|
"""
|
|
if tools is None and isinstance(toolcall_session, list):
|
|
tools, toolcall_session = toolcall_session, None
|
|
|
|
if tools and toolcall_session is None and all(is_tool(t) for t in tools):
|
|
session = FunctionToolSession(tools)
|
|
self.is_tools = True
|
|
self.toolcall_session = session
|
|
self.tools = session.schemas
|
|
return
|
|
|
|
if not (toolcall_session and tools):
|
|
return
|
|
self.is_tools = True
|
|
self.toolcall_session = toolcall_session
|
|
self.tools = tools
|
|
|
|
async def async_chat_with_tools(self, system: str, history: list, gen_conf: dict | None = None):
|
|
gen_conf = dict(gen_conf or {})
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
gen_conf, extra_request_kwargs = _apply_model_family_policies(
|
|
self.model_name,
|
|
backend="base",
|
|
gen_conf=gen_conf,
|
|
request_kwargs={},
|
|
)
|
|
if system and history and history[0].get("role") != "system":
|
|
history.insert(0, {"role": "system", "content": system})
|
|
|
|
ans = ""
|
|
tk_count = 0
|
|
# Aggregate prompt/completion/total across all tool-calling rounds.
|
|
agg_usage = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
|
|
|
def _add_round_usage(resp):
|
|
nonlocal tk_count
|
|
u = usage_from_response(resp)
|
|
agg_usage["prompt_tokens"] += u["prompt_tokens"]
|
|
agg_usage["completion_tokens"] += u["completion_tokens"]
|
|
agg_usage["total_tokens"] += u["total_tokens"] or total_token_count_from_response(resp)
|
|
tk_count = agg_usage["total_tokens"]
|
|
self.last_usage = dict(agg_usage)
|
|
|
|
hist = deepcopy(history)
|
|
for attempt in range(self.max_retries + 1):
|
|
history = deepcopy(hist)
|
|
try:
|
|
for _ in range(self.max_rounds + 1):
|
|
logging.info(f"{self.tools=}")
|
|
response = await self.async_client.chat.completions.create(model=self.model_name, messages=history, tools=self.tools, tool_choice="auto", **gen_conf, **extra_request_kwargs)
|
|
_add_round_usage(response)
|
|
if not response.choices or not response.choices[0].message:
|
|
raise Exception(f"500 response structure error. Response: {response}")
|
|
|
|
if not hasattr(response.choices[0].message, "tool_calls") or not response.choices[0].message.tool_calls:
|
|
_reasoning = getattr(response.choices[0].message, "reasoning_content", None) or getattr(response.choices[0].message, "reasoning", None)
|
|
if _reasoning:
|
|
ans += "<think>" + _reasoning + "</think>"
|
|
|
|
ans += response.choices[0].message.content
|
|
if response.choices[0].finish_reason == "length":
|
|
ans = self._length_stop(ans)
|
|
|
|
return ans, tk_count
|
|
|
|
async def _exec_tool(tc):
|
|
name = tc.function.name
|
|
try:
|
|
args = json_repair.loads(tc.function.arguments)
|
|
if not isinstance(args, dict):
|
|
raise TypeError(f"Tool arguments for {name} must be a JSON object, got {type(args).__name__}")
|
|
if hasattr(self.toolcall_session, "tool_call_async"):
|
|
result = await self.toolcall_session.tool_call_async(name, args)
|
|
else:
|
|
result = await thread_pool_exec(self.toolcall_session.tool_call, name, args)
|
|
return tc, name, args, result, None
|
|
except Exception as e:
|
|
logging.exception(f"Tool call failed: {tc}")
|
|
return tc, name, {}, None, e
|
|
|
|
logging.info(f"Response tool_calls={response.choices[0].message.tool_calls}")
|
|
results = await asyncio.gather(*[_exec_tool(tc) for tc in response.choices[0].message.tool_calls])
|
|
history = self._append_history_batch(history, results)
|
|
for tc, name, args, result, err in results:
|
|
ans += self._verbose_tool_use(name, args, err if err else result)
|
|
|
|
logging.warning(f"Exceed max rounds: {self.max_rounds}")
|
|
history.append({"role": "user", "content": f"Exceed max rounds: {self.max_rounds}"})
|
|
response, token_count = await self._async_chat(history, gen_conf)
|
|
ans += response
|
|
# _async_chat set self.last_usage to its own call; fold it into the aggregate.
|
|
_fb = getattr(self, "last_usage", None) or {}
|
|
agg_usage["prompt_tokens"] += int(_fb.get("prompt_tokens", 0) or 0)
|
|
agg_usage["completion_tokens"] += int(_fb.get("completion_tokens", 0) or 0)
|
|
agg_usage["total_tokens"] += int(_fb.get("total_tokens", 0) or token_count)
|
|
tk_count = agg_usage["total_tokens"]
|
|
self.last_usage = dict(agg_usage)
|
|
return ans, tk_count
|
|
except Exception as e:
|
|
e = await self._exceptions_async(e, attempt)
|
|
if e:
|
|
return e, tk_count
|
|
|
|
assert False, "Shouldn't be here."
|
|
|
|
async def async_chat_streamly_with_tools(self, system: str, history: list, gen_conf: dict | None = None):
|
|
gen_conf = dict(gen_conf or {})
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
gen_conf, extra_request_kwargs = _apply_model_family_policies(
|
|
self.model_name,
|
|
backend="base",
|
|
gen_conf=gen_conf,
|
|
request_kwargs={},
|
|
)
|
|
tools = self.tools
|
|
if system and history and history[0].get("role") != "system":
|
|
history.insert(0, {"role": "system", "content": system})
|
|
|
|
total_tokens = 0
|
|
# Aggregate prompt/completion/total across all tool-calling rounds. The split is
|
|
# captured opportunistically when the provider reports usage on a chunk; otherwise
|
|
# only the (estimated) total accumulates.
|
|
agg_usage = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
|
hist = deepcopy(history)
|
|
|
|
def _commit_round(round_usage, round_estimate):
|
|
nonlocal total_tokens
|
|
if round_usage and round_usage["total_tokens"]:
|
|
agg_usage["prompt_tokens"] += round_usage["prompt_tokens"]
|
|
agg_usage["completion_tokens"] += round_usage["completion_tokens"]
|
|
agg_usage["total_tokens"] += round_usage["total_tokens"]
|
|
else:
|
|
agg_usage["total_tokens"] += round_estimate
|
|
total_tokens = agg_usage["total_tokens"]
|
|
self.last_usage = dict(agg_usage)
|
|
|
|
for attempt in range(self.max_retries + 1):
|
|
history = deepcopy(hist)
|
|
try:
|
|
for _round in range(self.max_rounds + 1):
|
|
reasoning_start = False
|
|
logging.info(f"[Tool loop] Deciding what to do next (step {_round + 1}); available tools: {', '.join(t['function']['name'] for t in tools)}")
|
|
|
|
response = await self.async_client.chat.completions.create(
|
|
model=self.model_name, messages=history, stream=True, tools=tools, tool_choice="auto", **gen_conf, **extra_request_kwargs
|
|
)
|
|
|
|
final_tool_calls = {}
|
|
answer = ""
|
|
round_estimate = 0
|
|
round_usage = None
|
|
|
|
async for resp in response:
|
|
_u = usage_from_response(resp)
|
|
if _u["total_tokens"]:
|
|
round_usage = _u
|
|
|
|
if not hasattr(resp, "choices") or not resp.choices:
|
|
continue
|
|
|
|
delta = resp.choices[0].delta
|
|
|
|
if hasattr(delta, "tool_calls") and delta.tool_calls:
|
|
for tool_call in delta.tool_calls:
|
|
index = tool_call.index
|
|
if index not in final_tool_calls:
|
|
if not tool_call.function.arguments:
|
|
tool_call.function.arguments = ""
|
|
final_tool_calls[index] = tool_call
|
|
else:
|
|
final_tool_calls[index].function.arguments += tool_call.function.arguments or ""
|
|
continue
|
|
|
|
if not hasattr(delta, "content") or delta.content is None:
|
|
delta.content = ""
|
|
|
|
_reasoning = getattr(delta, "reasoning_content", None) or getattr(delta, "reasoning", None)
|
|
if _reasoning:
|
|
ans = ""
|
|
if not reasoning_start:
|
|
reasoning_start = True
|
|
ans = "<think>"
|
|
ans += _reasoning + "</think>"
|
|
yield ans
|
|
else:
|
|
reasoning_start = False
|
|
answer += delta.content
|
|
yield delta.content
|
|
|
|
if not _u["total_tokens"]:
|
|
round_estimate += num_tokens_from_string(delta.content)
|
|
|
|
finish_reason = getattr(resp.choices[0], "finish_reason", "")
|
|
if finish_reason == "length":
|
|
yield self._length_stop("")
|
|
|
|
# Commit this round's tokens (each round is a separate provider
|
|
# request — accumulate, never overwrite).
|
|
_commit_round(round_usage, round_estimate)
|
|
|
|
if answer and not final_tool_calls:
|
|
logging.info(f"[Tool loop] Answering directly at step {_round + 1} — no tool needed.")
|
|
yield total_tokens
|
|
return
|
|
|
|
async def _exec_tool(tc):
|
|
name = tc.function.name
|
|
try:
|
|
args = json_repair.loads(tc.function.arguments)
|
|
if not isinstance(args, dict):
|
|
raise TypeError(f"Tool arguments for {name} must be a JSON object, got {type(args).__name__}")
|
|
if hasattr(self.toolcall_session, "tool_call_async"):
|
|
result = await self.toolcall_session.tool_call_async(name, args)
|
|
else:
|
|
result = await thread_pool_exec(self.toolcall_session.tool_call, name, args)
|
|
return tc, name, args, result, None
|
|
except Exception as e:
|
|
logging.exception(f"Tool call failed: {tc}")
|
|
return tc, name, {}, None, e
|
|
|
|
tcs = list(final_tool_calls.values())
|
|
logging.info(f"[Tool loop] Step {_round + 1}: running {', '.join(tc.function.name for tc in tcs)}...")
|
|
for tc in tcs:
|
|
try:
|
|
args = json_repair.loads(tc.function.arguments)
|
|
except Exception:
|
|
args = {}
|
|
yield f"<think>Running the {tc.function.name} tool...</think>"
|
|
results = await asyncio.gather(*[_exec_tool(tc) for tc in tcs])
|
|
|
|
# Terminal-tool short-circuit: stream a terminal tool's
|
|
# result (already the final answer) and stop the loop.
|
|
_terminal = getattr(self, "terminal_tools", None)
|
|
if _terminal:
|
|
for tc, name, args, result, err in results:
|
|
if name in _terminal and not err:
|
|
logging.info(f"[Tool loop] The {name} tool produced the final answer — done.")
|
|
out = result if isinstance(result, str) else json.dumps(result, ensure_ascii=False)
|
|
if out:
|
|
yield out
|
|
yield total_tokens
|
|
return
|
|
|
|
history = self._append_history_batch(history, results)
|
|
for tc, name, args, result, err in results:
|
|
yield self._verbose_tool_use(name, args, err if err else result)
|
|
|
|
logging.warning(f"Exceed max rounds: {self.max_rounds}")
|
|
history.append({"role": "user", "content": f"Exceed max rounds: {self.max_rounds}"})
|
|
|
|
response = await self.async_client.chat.completions.create(
|
|
model=self.model_name,
|
|
messages=history,
|
|
stream=True,
|
|
tools=tools,
|
|
tool_choice="auto",
|
|
**gen_conf,
|
|
**extra_request_kwargs,
|
|
)
|
|
|
|
fb_estimate = 0
|
|
fb_usage = None
|
|
async for resp in response:
|
|
_u = usage_from_response(resp)
|
|
if _u["total_tokens"]:
|
|
fb_usage = _u
|
|
if not hasattr(resp, "choices") or not resp.choices:
|
|
continue
|
|
delta = resp.choices[0].delta
|
|
if not hasattr(delta, "content") or delta.content is None:
|
|
continue
|
|
if not _u["total_tokens"]:
|
|
fb_estimate += num_tokens_from_string(delta.content)
|
|
yield delta.content
|
|
|
|
_commit_round(fb_usage, fb_estimate)
|
|
yield total_tokens
|
|
return
|
|
|
|
except Exception as e:
|
|
e = await self._exceptions_async(e, attempt)
|
|
if e:
|
|
logging.error(f"async_chat_streamly failed: {e}")
|
|
yield e
|
|
yield total_tokens
|
|
return
|
|
|
|
assert False, "Shouldn't be here."
|
|
|
|
async def _async_chat(self, history, gen_conf, **kwargs):
|
|
logging.info("[HISTORY]" + json.dumps(history, ensure_ascii=False, indent=2))
|
|
if self.model_name.lower().find("qwq") >= 0:
|
|
logging.info(f"[INFO] {self.model_name} detected as reasoning model, using async_chat_streamly")
|
|
|
|
final_ans = ""
|
|
tol_token = 0
|
|
async for delta, tol in self._async_chat_streamly(history, gen_conf, with_reasoning=False, **kwargs):
|
|
if delta.startswith("<think>") or delta.endswith("</think>"):
|
|
continue
|
|
final_ans += delta
|
|
tol_token = tol
|
|
|
|
if len(final_ans.strip()) == 0:
|
|
final_ans = "**ERROR**: Empty response from reasoning model"
|
|
|
|
return final_ans.strip(), tol_token
|
|
|
|
gen_conf, kwargs = _apply_model_family_policies(
|
|
self.model_name,
|
|
backend="base",
|
|
gen_conf=gen_conf,
|
|
request_kwargs=kwargs,
|
|
)
|
|
|
|
response = await self.async_client.chat.completions.create(model=self.model_name, messages=history, **gen_conf, **kwargs)
|
|
|
|
# Capture prompt/completion split for accurate Langfuse + run aggregation.
|
|
self.last_usage = usage_from_response(response)
|
|
if not response.choices or not response.choices[0].message or not response.choices[0].message.content:
|
|
return "", 0
|
|
ans = response.choices[0].message.content.strip()
|
|
if response.choices[0].finish_reason == "length":
|
|
ans = self._length_stop(ans)
|
|
return ans, total_token_count_from_response(response)
|
|
|
|
async def async_chat(self, system, history, gen_conf=None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
if system and history and history[0].get("role") != "system":
|
|
history.insert(0, {"role": "system", "content": system})
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
|
|
for attempt in range(self.max_retries + 1):
|
|
try:
|
|
return await self._async_chat(history, gen_conf, **kwargs)
|
|
except Exception as e:
|
|
e = await self._exceptions_async(e, attempt)
|
|
if e:
|
|
return e, 0
|
|
assert False, "Shouldn't be here."
|
|
|
|
|
|
class XinferenceChat(Base):
|
|
_FACTORY_NAME = "Xinference"
|
|
|
|
def __init__(self, key=None, model_name="", base_url="", **kwargs):
|
|
if not base_url:
|
|
raise ValueError("Local llm url cannot be None")
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
|
|
class HuggingFaceChat(Base):
|
|
_FACTORY_NAME = "HuggingFace"
|
|
|
|
def __init__(self, key=None, model_name="", base_url="", **kwargs):
|
|
if not base_url:
|
|
raise ValueError("Local llm url cannot be None")
|
|
super().__init__(key, model_name.split("___")[0], base_url, **kwargs)
|
|
|
|
|
|
class ModelScopeChat(Base):
|
|
_FACTORY_NAME = "ModelScope"
|
|
|
|
def __init__(self, key=None, model_name="", base_url="", **kwargs):
|
|
if not base_url:
|
|
raise ValueError("Local llm url cannot be None")
|
|
super().__init__(key, model_name.split("___")[0], base_url, **kwargs)
|
|
|
|
|
|
class BaiChuanChat(Base):
|
|
_FACTORY_NAME = "BaiChuan"
|
|
|
|
def __init__(self, key, model_name="Baichuan3-Turbo", base_url="https://api.baichuan-ai.com/v1", **kwargs):
|
|
if not base_url:
|
|
base_url = "https://api.baichuan-ai.com/v1"
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
@staticmethod
|
|
def _format_params(params):
|
|
return {
|
|
"temperature": params.get("temperature", 0.3),
|
|
"top_p": params.get("top_p", 0.85),
|
|
}
|
|
|
|
def _clean_conf(self, gen_conf):
|
|
return {
|
|
"temperature": gen_conf.get("temperature", 0.3),
|
|
"top_p": gen_conf.get("top_p", 0.85),
|
|
}
|
|
|
|
def _chat(self, history, gen_conf=None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
response = self.client.chat.completions.create(
|
|
model=self.model_name,
|
|
messages=history,
|
|
extra_body={"tools": [{"type": "web_search", "web_search": {"enable": True, "search_mode": "performance_first"}}]},
|
|
**gen_conf,
|
|
)
|
|
if not response.choices:
|
|
raise ValueError("LLM returned empty response") # pact: guard empty choices list
|
|
ans = response.choices[0].message.content.strip()
|
|
if response.choices[0].finish_reason == "length":
|
|
if is_chinese([ans]):
|
|
ans += LENGTH_NOTIFICATION_CN
|
|
else:
|
|
ans += LENGTH_NOTIFICATION_EN
|
|
return ans, total_token_count_from_response(response)
|
|
|
|
def chat_streamly(self, system, history, gen_conf=None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
if system and history and history[0].get("role") != "system":
|
|
history.insert(0, {"role": "system", "content": system})
|
|
if "max_tokens" in gen_conf:
|
|
del gen_conf["max_tokens"]
|
|
ans = ""
|
|
total_tokens = 0
|
|
try:
|
|
response = self.client.chat.completions.create(
|
|
model=self.model_name,
|
|
messages=history,
|
|
extra_body={"tools": [{"type": "web_search", "web_search": {"enable": True, "search_mode": "performance_first"}}]},
|
|
stream=True,
|
|
**self._format_params(gen_conf),
|
|
)
|
|
for resp in response:
|
|
if not resp.choices:
|
|
continue
|
|
if not resp.choices[0].delta.content:
|
|
resp.choices[0].delta.content = ""
|
|
ans = resp.choices[0].delta.content
|
|
tol = total_token_count_from_response(resp)
|
|
if not tol:
|
|
total_tokens += num_tokens_from_string(resp.choices[0].delta.content)
|
|
else:
|
|
total_tokens = tol
|
|
if resp.choices[0].finish_reason == "length":
|
|
if is_chinese([ans]):
|
|
ans += LENGTH_NOTIFICATION_CN
|
|
else:
|
|
ans += LENGTH_NOTIFICATION_EN
|
|
yield ans
|
|
|
|
except Exception as e:
|
|
yield ans + "\n**ERROR**: " + str(e)
|
|
|
|
yield total_tokens
|
|
|
|
|
|
class LocalAIChat(Base):
|
|
_FACTORY_NAME = "LocalAI"
|
|
|
|
def __init__(self, key, model_name, base_url=None, **kwargs):
|
|
super().__init__(key, model_name, base_url=base_url, **kwargs)
|
|
|
|
if not base_url:
|
|
raise ValueError("Local llm url cannot be None")
|
|
self.client = OpenAI(api_key="empty", base_url=self.base_url)
|
|
self.model_name = model_name.split("___")[0]
|
|
|
|
|
|
class LocalLLM(Base):
|
|
def __init__(self, key, model_name, base_url=None, **kwargs):
|
|
super().__init__(key, model_name, base_url=base_url, **kwargs)
|
|
from jina import Client
|
|
|
|
self.client = Client(port=12345, protocol="grpc", asyncio=True)
|
|
|
|
def _prepare_prompt(self, system, history, gen_conf):
|
|
from rag.svr.jina_server import Prompt
|
|
|
|
if system and history and history[0].get("role") != "system":
|
|
history.insert(0, {"role": "system", "content": system})
|
|
return Prompt(message=history, gen_conf=gen_conf)
|
|
|
|
def _stream_response(self, endpoint, prompt):
|
|
from rag.svr.jina_server import Generation
|
|
|
|
answer = ""
|
|
loop = asyncio.new_event_loop()
|
|
try:
|
|
res = self.client.stream_doc(on=endpoint, inputs=prompt, return_type=Generation)
|
|
try:
|
|
while True:
|
|
answer = loop.run_until_complete(res.__anext__()).text
|
|
yield answer
|
|
except StopAsyncIteration:
|
|
pass
|
|
except Exception as e:
|
|
yield answer + "\n**ERROR**: " + str(e)
|
|
finally:
|
|
loop.close()
|
|
yield num_tokens_from_string(answer)
|
|
|
|
def chat(self, system, history, gen_conf=None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
if "max_tokens" in gen_conf:
|
|
del gen_conf["max_tokens"]
|
|
prompt = self._prepare_prompt(system, history, gen_conf)
|
|
chat_gen = self._stream_response("/chat", prompt)
|
|
ans = next(chat_gen)
|
|
total_tokens = next(chat_gen)
|
|
return ans, total_tokens
|
|
|
|
def chat_streamly(self, system, history, gen_conf=None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
if "max_tokens" in gen_conf:
|
|
del gen_conf["max_tokens"]
|
|
prompt = self._prepare_prompt(system, history, gen_conf)
|
|
return self._stream_response("/stream", prompt)
|
|
|
|
|
|
class VolcEngineChat(Base):
|
|
_FACTORY_NAME = "VolcEngine"
|
|
|
|
def __init__(self, key, model_name, base_url="https://ark.cn-beijing.volces.com/api/v3", **kwargs):
|
|
"""
|
|
Since do not want to modify the original database fields, and the VolcEngine authentication method is quite special,
|
|
Assemble ark_api_key, ep_id into api_key, store it as a dictionary type, and parse it for use
|
|
model_name is for display only
|
|
"""
|
|
base_url = base_url if base_url else "https://ark.cn-beijing.volces.com/api/v3"
|
|
try:
|
|
ark_api_key = json.loads(key).get("ark_api_key", "")
|
|
model_name = json.loads(key).get("ep_id", "") + json.loads(key).get("endpoint_id", "")
|
|
super().__init__(ark_api_key, model_name, base_url, **kwargs)
|
|
except JSONDecodeError:
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
|
|
class MistralChat(Base):
|
|
_FACTORY_NAME = "Mistral"
|
|
|
|
def __init__(self, key, model_name, base_url=None, **kwargs):
|
|
super().__init__(key, model_name, base_url=base_url, **kwargs)
|
|
|
|
from mistralai.client import Mistral
|
|
|
|
self.client = Mistral(api_key=key)
|
|
self.model_name = model_name
|
|
|
|
def _clean_conf(self, gen_conf):
|
|
for k in list(gen_conf.keys()):
|
|
if k not in ["temperature", "top_p", "max_tokens"]:
|
|
del gen_conf[k]
|
|
return gen_conf
|
|
|
|
def _chat(self, history, gen_conf=None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
response = self.client.chat.complete(model=self.model_name, messages=history, **gen_conf)
|
|
if not response.choices:
|
|
raise ValueError("LLM returned empty response") # pact: guard empty choices list
|
|
ans = response.choices[0].message.content
|
|
if response.choices[0].finish_reason == "length":
|
|
if is_chinese(ans):
|
|
ans += LENGTH_NOTIFICATION_CN
|
|
else:
|
|
ans += LENGTH_NOTIFICATION_EN
|
|
return ans, total_token_count_from_response(response)
|
|
|
|
def chat_streamly(self, system, history, gen_conf=None, **kwargs):
|
|
from mistralai.client.errors import MistralError
|
|
|
|
gen_conf = dict(gen_conf or {})
|
|
if system and history and history[0].get("role") != "system":
|
|
history.insert(0, {"role": "system", "content": system})
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
ans = ""
|
|
total_tokens = 0
|
|
try:
|
|
response = self.client.chat.stream(model=self.model_name, messages=history, **gen_conf, **kwargs)
|
|
for event in response:
|
|
resp = event.data
|
|
if not resp.choices or not resp.choices[0].delta.content:
|
|
continue
|
|
ans = resp.choices[0].delta.content
|
|
total_tokens += 1
|
|
if resp.choices[0].finish_reason == "length":
|
|
if is_chinese(ans):
|
|
ans += LENGTH_NOTIFICATION_CN
|
|
else:
|
|
ans += LENGTH_NOTIFICATION_EN
|
|
yield ans
|
|
|
|
except MistralError as e:
|
|
yield ans + "\n**ERROR**: " + str(e)
|
|
|
|
yield total_tokens
|
|
|
|
|
|
class LmStudioChat(Base):
|
|
_FACTORY_NAME = "LM-Studio"
|
|
|
|
def __init__(self, key, model_name, base_url, **kwargs):
|
|
if not base_url:
|
|
raise ValueError("Local llm url cannot be None")
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
self.client = OpenAI(api_key="lm-studio", base_url=self.base_url)
|
|
self.model_name = model_name
|
|
|
|
|
|
class OpenAI_APIChat(Base):
|
|
_FACTORY_NAME = ["VLLM", "OpenAI-API-Compatible"]
|
|
|
|
def __init__(self, key, model_name, base_url, **kwargs):
|
|
if not base_url:
|
|
raise ValueError("url cannot be None")
|
|
model_name = model_name.split("___")[0]
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
|
|
class MWSChat(Base):
|
|
"""MWS Chat Completions adapter with a documentation-only request body."""
|
|
|
|
_FACTORY_NAME = "MWS"
|
|
_ROLES = {"system", "user", "assistant"}
|
|
|
|
def __init__(self, key, model_name, base_url, **kwargs):
|
|
"""Initialize chat access for an MWS project and model deployment."""
|
|
token = require_mws_token(key)
|
|
self.chat_url = mws_api_url(base_url, "openai/v1/chat/completions")
|
|
self.headers = {
|
|
"Content-Type": "application/json",
|
|
"Authorization": f"Bearer {token}",
|
|
}
|
|
super().__init__(
|
|
token,
|
|
model_name.split("___")[0],
|
|
mws_api_url(base_url, "openai/v1"),
|
|
**kwargs,
|
|
)
|
|
|
|
def _clean_conf(self, gen_conf):
|
|
"""Keep only generation parameters documented by the MWS API."""
|
|
gen_conf = gen_conf or {}
|
|
cleaned = {}
|
|
if gen_conf.get("temperature") is not None:
|
|
cleaned["temperature"] = gen_conf["temperature"]
|
|
max_tokens = gen_conf.get("max_completion_tokens")
|
|
if max_tokens is None:
|
|
max_tokens = gen_conf.get("max_tokens")
|
|
if max_tokens is not None:
|
|
cleaned["max_completion_tokens"] = max_tokens
|
|
return cleaned
|
|
|
|
def _request_body(self, history, gen_conf, *, stream):
|
|
"""Build a strict MWS chat request from RAGFlow messages and options."""
|
|
messages = []
|
|
for message in history:
|
|
role = message.get("role") if isinstance(message, dict) else None
|
|
content = message.get("content") if isinstance(message, dict) else None
|
|
if role not in self._ROLES or not isinstance(content, str):
|
|
raise ValueError("MWS chat messages must contain only a system, user, or assistant role and string content")
|
|
messages.append({"role": role, "content": content})
|
|
if not messages:
|
|
raise ValueError("MWS chat messages are required")
|
|
|
|
body = {"model": self.model_name, "messages": messages}
|
|
body.update(self._clean_conf(gen_conf))
|
|
if stream:
|
|
body["stream"] = True
|
|
body["stream_options"] = {"include_usage": True}
|
|
return body
|
|
|
|
async def _post_json(self, body):
|
|
"""Send a non-streaming MWS chat request and decode its JSON response."""
|
|
timeout = aiohttp.ClientTimeout(total=int(os.environ.get("LLM_TIMEOUT_SECONDS", 600)))
|
|
async with aiohttp.ClientSession(timeout=timeout) as session:
|
|
async with session.post(
|
|
self.chat_url,
|
|
headers=self.headers,
|
|
json=body,
|
|
) as response:
|
|
if response.status != 200:
|
|
raise RuntimeError(f"MWS chat request failed with status {response.status}: {await response.text()}")
|
|
return await response.json()
|
|
|
|
async def _async_chat(self, history, gen_conf, **kwargs):
|
|
"""Return one complete MWS chat answer together with its token usage."""
|
|
payload = await self._post_json(self._request_body(history, gen_conf, stream=False))
|
|
self.last_usage = usage_from_response(payload)
|
|
choices = payload.get("choices") if isinstance(payload, dict) else None
|
|
if not isinstance(choices, list) or not choices:
|
|
raise ValueError("MWS chat response does not contain choices")
|
|
choice = choices[0]
|
|
message = choice.get("message") if isinstance(choice, dict) else None
|
|
content = message.get("content") if isinstance(message, dict) else None
|
|
if not isinstance(content, str):
|
|
raise ValueError("MWS chat response does not contain message content")
|
|
answer = content.strip()
|
|
if choice.get("finish_reason") == "length":
|
|
answer = self._length_stop(answer)
|
|
return answer, total_token_count_from_response(payload)
|
|
|
|
async def _async_chat_streamly(self, history, gen_conf, **kwargs):
|
|
"""Yield MWS SSE content chunks and attach usage to the final chunk."""
|
|
body = self._request_body(history, gen_conf, stream=True)
|
|
timeout = aiohttp.ClientTimeout(total=int(os.environ.get("LLM_TIMEOUT_SECONDS", 600)))
|
|
pending_content = None
|
|
estimated_tokens = 0
|
|
reported_tokens = 0
|
|
|
|
async with aiohttp.ClientSession(timeout=timeout) as session:
|
|
async with session.post(
|
|
self.chat_url,
|
|
headers=self.headers,
|
|
json=body,
|
|
) as response:
|
|
if response.status != 200:
|
|
raise RuntimeError(f"MWS chat request failed with status {response.status}: {await response.text()}")
|
|
|
|
async for raw_line in response.content:
|
|
line = raw_line.decode("utf-8").strip()
|
|
if not line.startswith("data:"):
|
|
continue
|
|
data = line[5:].strip()
|
|
if data == "[DONE]":
|
|
break
|
|
event = json.loads(data)
|
|
usage = usage_from_response(event)
|
|
if usage["total_tokens"]:
|
|
self.last_usage = usage
|
|
reported_tokens = usage["total_tokens"]
|
|
|
|
choices = event.get("choices")
|
|
if not isinstance(choices, list) or not choices:
|
|
continue
|
|
choice = choices[0]
|
|
delta = choice.get("delta") if isinstance(choice, dict) else None
|
|
content = delta.get("content") if isinstance(delta, dict) else None
|
|
if not isinstance(content, str) or not content:
|
|
continue
|
|
if choice.get("finish_reason") == "length":
|
|
content = self._length_stop(content)
|
|
if pending_content is not None:
|
|
yield pending_content, 0
|
|
pending_content = content
|
|
estimated_tokens += num_tokens_from_string(content)
|
|
|
|
yield pending_content or "", reported_tokens or estimated_tokens
|
|
|
|
|
|
class Xiaomi(Base):
|
|
_FACTORY_NAME = "Xiaomi"
|
|
|
|
def __init__(self, key, model_name, base_url, **kwargs):
|
|
if not base_url:
|
|
base_url = "https://api.xiaomimimo.com/v1"
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
|
|
class LeptonAIChat(Base):
|
|
_FACTORY_NAME = "LeptonAI"
|
|
|
|
def __init__(self, key, model_name, base_url=None, **kwargs):
|
|
if not base_url:
|
|
base_url = urljoin("https://" + model_name + ".lepton.run", "api/v1")
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
|
|
class ReplicateChat(Base):
|
|
_FACTORY_NAME = "Replicate"
|
|
|
|
def __init__(self, key, model_name, base_url=None, **kwargs):
|
|
super().__init__(key, model_name, base_url=base_url, **kwargs)
|
|
|
|
from replicate.client import Client
|
|
|
|
self.model_name = model_name
|
|
self.client = Client(api_token=_normalize_replicate_key(key))
|
|
|
|
def _chat(self, history, gen_conf=None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
system = history[0]["content"] if history and history[0]["role"] == "system" else ""
|
|
prompt = "\n".join([item["role"] + ":" + item["content"] for item in history[-5:] if item["role"] != "system"])
|
|
response = self.client.run(
|
|
self.model_name,
|
|
input={"system_prompt": system, "prompt": prompt, **gen_conf},
|
|
)
|
|
ans = "".join(response)
|
|
return ans, num_tokens_from_string(ans)
|
|
|
|
def chat_streamly(self, system, history, gen_conf=None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
if "max_tokens" in gen_conf:
|
|
del gen_conf["max_tokens"]
|
|
prompt = "\n".join([item["role"] + ":" + item["content"] for item in history[-5:]])
|
|
ans = ""
|
|
try:
|
|
response = self.client.run(
|
|
self.model_name,
|
|
input={"system_prompt": system, "prompt": prompt, **gen_conf},
|
|
)
|
|
for resp in response:
|
|
ans = resp
|
|
yield ans
|
|
|
|
except Exception as e:
|
|
yield ans + "\n**ERROR**: " + str(e)
|
|
|
|
yield num_tokens_from_string(ans)
|
|
|
|
async def async_chat_streamly(self, system, history, gen_conf: dict | None = None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
if "max_tokens" in gen_conf:
|
|
del gen_conf["max_tokens"]
|
|
|
|
def _do_chat():
|
|
msgs = list(history or [])
|
|
if system and msgs and msgs[0].get("role") != "system":
|
|
msgs.insert(0, {"role": "system", "content": system})
|
|
elif system and not msgs:
|
|
msgs = [{"role": "system", "content": system}]
|
|
|
|
system_msg = msgs[0]["content"] if msgs and msgs[0].get("role") == "system" else ""
|
|
prompt = "\n".join([item["role"] + ":" + item["content"] for item in msgs[-5:] if item.get("role") != "system"])
|
|
try:
|
|
response = self.client.run(
|
|
self.model_name,
|
|
input={"system_prompt": system_msg, "prompt": prompt, **gen_conf},
|
|
)
|
|
chunks = []
|
|
for resp in response:
|
|
chunks.append(resp if isinstance(resp, str) else str(resp))
|
|
answer = "".join(chunks)
|
|
return chunks or ([answer] if answer else []), num_tokens_from_string(answer), None
|
|
except Exception as e:
|
|
return [], 0, e
|
|
|
|
chunks, total_tokens, error = await asyncio.to_thread(_do_chat)
|
|
if error:
|
|
yield f"**ERROR**: {error}"
|
|
else:
|
|
for chunk in chunks:
|
|
yield chunk
|
|
yield total_tokens
|
|
|
|
|
|
class SparkChat(Base):
|
|
_FACTORY_NAME = "XunFei Spark"
|
|
|
|
def __init__(self, key, model_name, base_url="https://spark-api-open.xf-yun.com/v1", **kwargs):
|
|
if not base_url:
|
|
base_url = "https://spark-api-open.xf-yun.com/v1"
|
|
model2version = {
|
|
"Spark-Max": "generalv3.5",
|
|
"Spark-Max-32K": "max-32k",
|
|
"Spark-Lite": "lite",
|
|
"Spark-Pro": "generalv3",
|
|
"Spark-Pro-128K": "pro-128k",
|
|
"Spark-4.0-Ultra": "4.0Ultra",
|
|
}
|
|
version2model = {v: k for k, v in model2version.items()}
|
|
assert model_name in model2version or model_name in version2model, f"The given model name is not supported yet. Support: {list(model2version.keys())}"
|
|
if model_name in model2version:
|
|
model_version = model2version[model_name]
|
|
else:
|
|
model_version = model_name
|
|
super().__init__(key, model_version, base_url, **kwargs)
|
|
|
|
|
|
class BaiduYiyanChat(Base):
|
|
_FACTORY_NAME = "BaiduYiyan"
|
|
|
|
def __init__(self, key, model_name, base_url=None, **kwargs):
|
|
super().__init__(key, model_name, base_url=base_url, **kwargs)
|
|
|
|
import qianfan
|
|
|
|
key = json.loads(key)
|
|
ak = key.get("yiyan_ak", "")
|
|
sk = key.get("yiyan_sk", "")
|
|
self.client = qianfan.ChatCompletion(ak=ak, sk=sk)
|
|
self.model_name = model_name.lower()
|
|
|
|
def _clean_conf(self, gen_conf):
|
|
gen_conf["penalty_score"] = ((gen_conf.get("presence_penalty", 0) + gen_conf.get("frequency_penalty", 0)) / 2) + 1
|
|
if "max_tokens" in gen_conf:
|
|
del gen_conf["max_tokens"]
|
|
return gen_conf
|
|
|
|
def _chat(self, history, gen_conf):
|
|
system = history[0]["content"] if history and history[0]["role"] == "system" else ""
|
|
response = self.client.do(model=self.model_name, messages=[h for h in history if h["role"] != "system"], system=system, **gen_conf).body
|
|
ans = response["result"]
|
|
return ans, total_token_count_from_response(response)
|
|
|
|
def chat_streamly(self, system, history, gen_conf=None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
gen_conf["penalty_score"] = ((gen_conf.get("presence_penalty", 0) + gen_conf.get("frequency_penalty", 0)) / 2) + 1
|
|
if "max_tokens" in gen_conf:
|
|
del gen_conf["max_tokens"]
|
|
ans = ""
|
|
total_tokens = 0
|
|
|
|
try:
|
|
response = self.client.do(model=self.model_name, messages=history, system=system, stream=True, **gen_conf)
|
|
for resp in response:
|
|
resp = resp.body
|
|
ans = resp["result"]
|
|
total_tokens = total_token_count_from_response(resp)
|
|
|
|
yield ans
|
|
|
|
except Exception as e:
|
|
return ans + "\n**ERROR**: " + str(e), 0
|
|
|
|
yield total_tokens
|
|
|
|
async def async_chat_streamly(self, system, history, gen_conf: dict | None = None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
gen_conf["penalty_score"] = ((gen_conf.get("presence_penalty", 0) + gen_conf.get("frequency_penalty", 0)) / 2) + 1
|
|
if "max_tokens" in gen_conf:
|
|
del gen_conf["max_tokens"]
|
|
|
|
def _do_chat():
|
|
system_msg = history[0]["content"] if history and history[0].get("role") == "system" else ""
|
|
msgs = [h for h in history if h.get("role") != "system"]
|
|
try:
|
|
response = self.client.do(model=self.model_name, messages=msgs, system=system_msg, stream=True, **gen_conf)
|
|
result_text = ""
|
|
total_tokens = 0
|
|
for resp in response:
|
|
resp = resp.body
|
|
result_text = resp["result"]
|
|
total_tokens = total_token_count_from_response(resp)
|
|
return result_text, total_tokens, None
|
|
except Exception as e:
|
|
return "", 0, e
|
|
|
|
result_text, total_tokens, error = await asyncio.to_thread(_do_chat)
|
|
if error:
|
|
yield f"**ERROR**: {error}"
|
|
else:
|
|
yield result_text
|
|
yield total_tokens
|
|
|
|
|
|
class GoogleChat(Base):
|
|
_FACTORY_NAME = "Google Cloud"
|
|
|
|
@staticmethod
|
|
def _vertex_http_options(region: str):
|
|
region_norm = (region or "").strip().lower()
|
|
multipoint_hosts = {
|
|
"eu": "https://aiplatform.eu.rep.googleapis.com/",
|
|
"us": "https://aiplatform.us.rep.googleapis.com/",
|
|
}
|
|
base_url = multipoint_hosts.get(region_norm)
|
|
if base_url:
|
|
from google.genai.types import HttpOptions
|
|
|
|
# Gemini 3.x multi-region endpoints require *.rep hostnames
|
|
# instead of region-aiplatform host synthesis.
|
|
return HttpOptions(base_url=base_url, api_version="v1")
|
|
return None
|
|
|
|
def __init__(self, key, model_name, base_url=None, **kwargs):
|
|
super().__init__(key, model_name, base_url=base_url, **kwargs)
|
|
|
|
import base64
|
|
|
|
from google.oauth2 import service_account
|
|
|
|
key = json.loads(key)
|
|
access_token = json.loads(base64.b64decode(key.get("google_service_account_key", "")))
|
|
project_id = key.get("google_project_id", "")
|
|
region = key.get("google_region", "")
|
|
|
|
scopes = ["https://www.googleapis.com/auth/cloud-platform"]
|
|
self.model_name = model_name
|
|
|
|
if "claude" in self.model_name:
|
|
from anthropic import AnthropicVertex
|
|
from google.auth.transport.requests import Request
|
|
|
|
if access_token:
|
|
credits = service_account.Credentials.from_service_account_info(access_token, scopes=scopes)
|
|
request = Request()
|
|
credits.refresh(request)
|
|
token = credits.token
|
|
self.client = AnthropicVertex(region=region, project_id=project_id, access_token=token)
|
|
else:
|
|
self.client = AnthropicVertex(region=region, project_id=project_id)
|
|
else:
|
|
from google import genai
|
|
|
|
client_kwargs = {
|
|
"vertexai": True,
|
|
"project": project_id,
|
|
"location": region,
|
|
}
|
|
http_options = self._vertex_http_options(region)
|
|
if http_options is not None:
|
|
client_kwargs["http_options"] = http_options
|
|
|
|
if access_token:
|
|
credits = service_account.Credentials.from_service_account_info(access_token, scopes=scopes)
|
|
client_kwargs["credentials"] = credits
|
|
self.client = genai.Client(**client_kwargs)
|
|
else:
|
|
self.client = genai.Client(**client_kwargs)
|
|
|
|
def _clean_conf(self, gen_conf):
|
|
if "claude" in self.model_name:
|
|
if "max_tokens" in gen_conf:
|
|
del gen_conf["max_tokens"]
|
|
else:
|
|
if "max_tokens" in gen_conf:
|
|
gen_conf["max_output_tokens"] = gen_conf["max_tokens"]
|
|
del gen_conf["max_tokens"]
|
|
for k in list(gen_conf.keys()):
|
|
if k not in ["temperature", "top_p", "max_output_tokens"]:
|
|
del gen_conf[k]
|
|
return gen_conf
|
|
|
|
def _chat(self, history, gen_conf=None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
system = history[0]["content"] if history and history[0]["role"] == "system" else ""
|
|
|
|
if "claude" in self.model_name:
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
response = self.client.messages.create(
|
|
model=self.model_name,
|
|
messages=[h for h in history if h["role"] != "system"],
|
|
system=system,
|
|
stream=False,
|
|
**gen_conf,
|
|
).json()
|
|
ans = response["content"][0]["text"]
|
|
if response["stop_reason"] == "max_tokens":
|
|
ans += "...\nFor the content length reason, it stopped, continue?" if is_english([ans]) else "······\n由于长度的原因,回答被截断了,要继续吗?"
|
|
return (
|
|
ans,
|
|
response["usage"]["input_tokens"] + response["usage"]["output_tokens"],
|
|
)
|
|
|
|
# Gemini models with google-genai SDK
|
|
# Set default thinking_budget=0 if not specified
|
|
if "thinking_budget" not in gen_conf:
|
|
gen_conf["thinking_budget"] = 0
|
|
|
|
thinking_budget = gen_conf.pop("thinking_budget", 0)
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
|
|
# Build GenerateContentConfig
|
|
try:
|
|
from google.genai.types import Content, GenerateContentConfig, Part, ThinkingConfig
|
|
except ImportError as e:
|
|
logging.error(f"[GoogleChat] Failed to import google-genai: {e}. Please install: pip install google-genai>=1.41.0")
|
|
raise
|
|
|
|
config_dict = {}
|
|
if system:
|
|
config_dict["system_instruction"] = system
|
|
if "temperature" in gen_conf:
|
|
config_dict["temperature"] = gen_conf["temperature"]
|
|
if "top_p" in gen_conf:
|
|
config_dict["top_p"] = gen_conf["top_p"]
|
|
if "max_output_tokens" in gen_conf:
|
|
config_dict["max_output_tokens"] = gen_conf["max_output_tokens"]
|
|
|
|
# Add ThinkingConfig
|
|
config_dict["thinking_config"] = ThinkingConfig(thinking_budget=thinking_budget)
|
|
|
|
config = GenerateContentConfig(**config_dict)
|
|
|
|
# Convert history to google-genai Content format
|
|
contents = []
|
|
for item in history:
|
|
if item["role"] == "system":
|
|
continue
|
|
# google-genai uses 'model' instead of 'assistant'
|
|
role = "model" if item["role"] == "assistant" else item["role"]
|
|
content = Content(
|
|
role=role,
|
|
parts=[Part(text=item["content"])],
|
|
)
|
|
contents.append(content)
|
|
|
|
response = self.client.models.generate_content(
|
|
model=self.model_name,
|
|
contents=contents,
|
|
config=config,
|
|
)
|
|
|
|
ans = response.text
|
|
# Get token count from response
|
|
try:
|
|
total_tokens = response.usage_metadata.total_token_count
|
|
except Exception:
|
|
total_tokens = 0
|
|
|
|
return ans, total_tokens
|
|
|
|
def chat_streamly(self, system, history, gen_conf=None, **kwargs):
|
|
gen_conf = dict(gen_conf or {})
|
|
if "claude" in self.model_name:
|
|
if "max_tokens" in gen_conf:
|
|
del gen_conf["max_tokens"]
|
|
ans = ""
|
|
total_tokens = 0
|
|
try:
|
|
response = self.client.messages.create(
|
|
model=self.model_name,
|
|
messages=history,
|
|
system=system,
|
|
stream=True,
|
|
**gen_conf,
|
|
)
|
|
for res in response.iter_lines():
|
|
res = res.decode("utf-8")
|
|
if "content_block_delta" in res and "data" in res:
|
|
text = json.loads(res[6:])["delta"]["text"]
|
|
ans = text
|
|
total_tokens += num_tokens_from_string(text)
|
|
except Exception as e:
|
|
yield ans + "\n**ERROR**: " + str(e)
|
|
|
|
yield total_tokens
|
|
else:
|
|
# Gemini models with google-genai SDK
|
|
ans = ""
|
|
total_tokens = 0
|
|
|
|
# Set default thinking_budget=0 if not specified
|
|
if "thinking_budget" not in gen_conf:
|
|
gen_conf["thinking_budget"] = 0
|
|
|
|
thinking_budget = gen_conf.pop("thinking_budget", 0)
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
|
|
# Build GenerateContentConfig
|
|
try:
|
|
from google.genai.types import Content, GenerateContentConfig, Part, ThinkingConfig
|
|
except ImportError as e:
|
|
logging.error(f"[GoogleChat] Failed to import google-genai: {e}. Please install: pip install google-genai>=1.41.0")
|
|
raise
|
|
|
|
config_dict = {}
|
|
if system:
|
|
config_dict["system_instruction"] = system
|
|
if "temperature" in gen_conf:
|
|
config_dict["temperature"] = gen_conf["temperature"]
|
|
if "top_p" in gen_conf:
|
|
config_dict["top_p"] = gen_conf["top_p"]
|
|
if "max_output_tokens" in gen_conf:
|
|
config_dict["max_output_tokens"] = gen_conf["max_output_tokens"]
|
|
|
|
# Add ThinkingConfig
|
|
config_dict["thinking_config"] = ThinkingConfig(thinking_budget=thinking_budget)
|
|
|
|
config = GenerateContentConfig(**config_dict)
|
|
|
|
# Convert history to google-genai Content format
|
|
contents = []
|
|
for item in history:
|
|
# google-genai uses 'model' instead of 'assistant'
|
|
role = "model" if item["role"] == "assistant" else item["role"]
|
|
content = Content(
|
|
role=role,
|
|
parts=[Part(text=item["content"])],
|
|
)
|
|
contents.append(content)
|
|
|
|
try:
|
|
for chunk in self.client.models.generate_content_stream(
|
|
model=self.model_name,
|
|
contents=contents,
|
|
config=config,
|
|
):
|
|
text = chunk.text
|
|
ans = text
|
|
total_tokens += num_tokens_from_string(text)
|
|
yield ans
|
|
|
|
except Exception as e:
|
|
yield ans + "\n**ERROR**: " + str(e)
|
|
|
|
yield total_tokens
|
|
|
|
|
|
class TokenPonyChat(Base):
|
|
_FACTORY_NAME = "TokenPony"
|
|
|
|
def __init__(self, key, model_name, base_url="https://ragflow.vip-api.tokenpony.cn/v1", **kwargs):
|
|
if not base_url:
|
|
base_url = "https://ragflow.vip-api.tokenpony.cn/v1"
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
|
|
class N1nChat(Base):
|
|
_FACTORY_NAME = "n1n"
|
|
|
|
def __init__(self, key, model_name, base_url="https://api.n1n.ai/v1", **kwargs):
|
|
if not base_url:
|
|
base_url = "https://api.n1n.ai/v1"
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
|
|
class AvianChat(Base):
|
|
_FACTORY_NAME = "Avian"
|
|
|
|
def __init__(self, key, model_name, base_url="https://api.avian.io/v1", **kwargs):
|
|
if not base_url:
|
|
base_url = "https://api.avian.io/v1"
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
|
|
class AstraflowChat(Base):
|
|
_FACTORY_NAME = "Astraflow"
|
|
|
|
def __init__(self, key, model_name, base_url="https://api-us-ca.umodelverse.ai/v1", **kwargs):
|
|
if not base_url:
|
|
base_url = "https://api-us-ca.umodelverse.ai/v1"
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
|
|
class AstraflowCNChat(Base):
|
|
_FACTORY_NAME = "Astraflow-CN"
|
|
|
|
def __init__(self, key, model_name, base_url="https://api.modelverse.cn/v1", **kwargs):
|
|
if not base_url:
|
|
base_url = "https://api.modelverse.cn/v1"
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
|
|
class FuturMixChat(Base):
|
|
_FACTORY_NAME = "FuturMix"
|
|
|
|
def __init__(self, key, model_name, base_url="https://futurmix.ai/v1", **kwargs):
|
|
if not base_url:
|
|
base_url = "https://futurmix.ai/v1"
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
logging.info("[FuturMix] Chat initialized with model %s", model_name)
|
|
|
|
|
|
class AIMLAPIChat(Base):
|
|
_FACTORY_NAME = "aimlapi.com"
|
|
|
|
def __init__(self, key, model_name, base_url="", **kwargs):
|
|
base_url = base_url or os.environ.get("AIMLAPI_API_URL", "https://api.aimlapi.com/v1")
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
headers = attribution_headers()
|
|
self.client = self.client.with_options(default_headers=headers)
|
|
self.async_client = self.async_client.with_options(default_headers=headers)
|
|
logging.info("[aimlapi.com] Chat initialized with model %s", model_name)
|
|
|
|
|
|
class GreenPTChat(Base):
|
|
"""GreenPT OpenAI-compatible chat adapter."""
|
|
|
|
_FACTORY_NAME = "GreenPT"
|
|
|
|
def __init__(self, key, model_name, base_url="https://api.greenpt.ai/v1", **kwargs):
|
|
super().__init__(key, model_name, base_url or "https://api.greenpt.ai/v1", **kwargs)
|
|
|
|
|
|
# MiniMax models sometimes emit their bracket-delimited control/boundary tokens
|
|
# into `content` instead of as structured control — e.g. "]<]minimax[>[" — most
|
|
# often on tool-calling turns. The token is streamed split across many deltas,
|
|
# so it can't be removed per-delta; it must be filtered over a window that spans
|
|
# chunk boundaries. This pattern only matches the vendor name when it is wrapped
|
|
# in bracket noise on BOTH sides, so ordinary prose that mentions "MiniMax" is
|
|
# left untouched. Extend the alternation as further control tokens are observed.
|
|
_MINIMAX_CONTROL_TOKEN_RE = re.compile(r"[\[\]<>]+\s*minimax\s*[\[\]<>]+", re.IGNORECASE)
|
|
|
|
|
|
class _StreamSanitizer:
|
|
"""Strip a regex from a token stream even when matches span chunk boundaries.
|
|
|
|
A control token is bracket+letter characters, and it can arrive split across
|
|
many deltas, so we hold back the trailing run of token-ish characters (which
|
|
might still be forming a match) and only ``sub`` + emit the part before it.
|
|
Applying ``sub`` to a partial trailing run would fire prematurely and leak the
|
|
unmatched remainder — hence the hold. ``flush()`` sanitizes and returns the
|
|
remainder at end of stream. ``keep`` caps how long a run is buffered so a very
|
|
long separator-less word can't stall the stream forever.
|
|
"""
|
|
|
|
_TOKENISH = re.compile(r"[\[\]<>A-Za-z]*$")
|
|
|
|
def __init__(self, pattern: re.Pattern, keep: int = 64) -> None:
|
|
self._pat = pattern
|
|
self._keep = keep
|
|
self._buf = ""
|
|
|
|
def feed(self, text: str) -> str:
|
|
if not text:
|
|
return ""
|
|
self._buf += text
|
|
match = self._TOKENISH.search(self._buf)
|
|
hold_start = match.start() if match else len(self._buf)
|
|
if len(self._buf) - hold_start > self._keep:
|
|
hold_start = len(self._buf) - self._keep
|
|
emit = self._pat.sub("", self._buf[:hold_start])
|
|
self._buf = self._buf[hold_start:]
|
|
return emit
|
|
|
|
def flush(self) -> str:
|
|
out = self._pat.sub("", self._buf)
|
|
self._buf = ""
|
|
return out
|
|
|
|
|
|
class LiteLLMBase(ABC):
|
|
_FACTORY_NAME = [
|
|
"Tongyi-Qianwen",
|
|
"Bedrock",
|
|
"Moonshot",
|
|
"xAI",
|
|
"DeepInfra",
|
|
"Groq",
|
|
"Cohere",
|
|
"Gemini",
|
|
"DeepSeek",
|
|
"NVIDIA",
|
|
"TogetherAI",
|
|
"Anthropic",
|
|
"Ollama",
|
|
"LongCat",
|
|
"CometAPI",
|
|
"SILICONFLOW",
|
|
"OpenRouter",
|
|
"StepFun",
|
|
"PPIO",
|
|
"PerfXCloud",
|
|
"Upstage",
|
|
"NovitaAI",
|
|
"01.AI",
|
|
"GiteeAI",
|
|
"302.AI",
|
|
"Jiekou.AI",
|
|
"ZHIPU-AI",
|
|
"MiniMax",
|
|
"DeerAPI",
|
|
"GPUStack",
|
|
"OpenAI",
|
|
"Azure-OpenAI",
|
|
"Tencent Hunyuan",
|
|
]
|
|
|
|
def __init__(self, key, model_name, base_url=None, **kwargs):
|
|
self.timeout = int(os.environ.get("LLM_TIMEOUT_SECONDS", 600))
|
|
self.provider = kwargs.get("provider", "")
|
|
self.prefix = LITELLM_PROVIDER_PREFIX.get(self.provider, "")
|
|
self.model_name = f"{self.prefix}{model_name}"
|
|
self.api_key = key
|
|
self.base_url = (base_url or FACTORY_DEFAULT_BASE_URL.get(self.provider, "")).rstrip("/")
|
|
# Configure retry parameters
|
|
self.max_retries = kwargs.get("max_retries", int(os.environ.get("LLM_MAX_RETRIES", 5)))
|
|
self.base_delay = kwargs.get("retry_interval", float(os.environ.get("LLM_BASE_DELAY", 2.0)))
|
|
self.max_rounds = kwargs.get("max_rounds", 5)
|
|
self.is_tools = False
|
|
self.tools = []
|
|
self.toolcall_sessions = {}
|
|
# Token usage split (prompt/completion/total) of the most recent chat call.
|
|
# Consumed by LLMBundle for accurate Langfuse reporting and run aggregation.
|
|
self.last_usage = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
|
|
|
# Factory specific fields
|
|
if self.provider == SupportedLiteLLMProvider.OpenRouter:
|
|
try:
|
|
self.api_key = json.loads(key).get("api_key", "")
|
|
self.provider_order = json.loads(key).get("provider_order", "")
|
|
except JSONDecodeError:
|
|
self.api_key = key
|
|
self.provider_order = ""
|
|
elif self.provider == SupportedLiteLLMProvider.Azure_OpenAI:
|
|
self.api_key = json.loads(key).get("api_key", "")
|
|
self.api_version = json.loads(key).get("api_version", "2024-02-01")
|
|
elif self.provider == SupportedLiteLLMProvider.MiniMax:
|
|
# MiniMax requires GroupId as a query parameter for API authentication
|
|
try:
|
|
key_obj = json.loads(key) if isinstance(key, str) else key
|
|
self.api_key = key_obj.get("api_key", key) if isinstance(key_obj, dict) else key
|
|
self.group_id = key_obj.get("group_id", "") if isinstance(key_obj, dict) else ""
|
|
except (json.JSONDecodeError, TypeError):
|
|
self.api_key = key
|
|
self.group_id = ""
|
|
else:
|
|
self.group_id = ""
|
|
|
|
def _get_delay(self):
|
|
return self.base_delay * random.uniform(10, 150)
|
|
|
|
def _classify_error(self, error):
|
|
error_str = str(error).lower()
|
|
|
|
keywords_mapping = [
|
|
(["quota", "capacity", "credit", "billing", "balance", "欠费"], LLMErrorCode.ERROR_QUOTA),
|
|
(["rate limit", "429", "tpm limit", "too many requests", "requests per minute"], LLMErrorCode.ERROR_RATE_LIMIT),
|
|
(["auth", "key", "apikey", "401", "forbidden", "permission"], LLMErrorCode.ERROR_AUTHENTICATION),
|
|
(["invalid", "bad request", "400", "format", "malformed", "parameter"], LLMErrorCode.ERROR_INVALID_REQUEST),
|
|
(["server", "503", "502", "504", "500", "unavailable"], LLMErrorCode.ERROR_SERVER),
|
|
(["timeout", "timed out"], LLMErrorCode.ERROR_TIMEOUT),
|
|
(["connect", "network", "unreachable", "dns"], LLMErrorCode.ERROR_CONNECTION),
|
|
(["filter", "content", "policy", "blocked", "safety", "inappropriate"], LLMErrorCode.ERROR_CONTENT_FILTER),
|
|
(["model", "not found", "does not exist", "not available"], LLMErrorCode.ERROR_MODEL),
|
|
(["max rounds"], LLMErrorCode.ERROR_MODEL),
|
|
]
|
|
for words, code in keywords_mapping:
|
|
if re.search("({})".format("|".join(words)), error_str):
|
|
return code
|
|
|
|
return LLMErrorCode.ERROR_GENERIC
|
|
|
|
def _clean_conf(self, gen_conf):
|
|
gen_conf, _ = _apply_model_family_policies(
|
|
self.model_name,
|
|
backend="litellm",
|
|
provider=self.provider,
|
|
gen_conf=gen_conf,
|
|
)
|
|
|
|
deepseek_max_tokens = None
|
|
if self.provider == SupportedLiteLLMProvider.DeepSeek:
|
|
# DeepSeek's API uses the legacy OpenAI-compatible max_tokens field.
|
|
# LiteLLM accepts max_completion_tokens generically, but does not
|
|
# translate it for DeepSeek and the provider then falls back to 8192.
|
|
# Knowledge compilation supplies max_completion_tokens explicitly
|
|
# through its model-specific generation configuration. Otherwise,
|
|
# preserve the legacy max_tokens value used by existing callers.
|
|
raw_max_completion_tokens = gen_conf.pop("max_completion_tokens", None)
|
|
raw_max_tokens = gen_conf.pop("max_tokens", None)
|
|
raw_limit = raw_max_completion_tokens if raw_max_completion_tokens is not None else raw_max_tokens
|
|
if raw_limit is not None and not isinstance(raw_limit, bool):
|
|
try:
|
|
candidate = int(raw_limit)
|
|
except (TypeError, ValueError):
|
|
candidate = 0
|
|
if candidate > 0:
|
|
deepseek_max_tokens = candidate
|
|
else:
|
|
gen_conf.pop("max_tokens", None)
|
|
|
|
gen_conf = {k: v for k, v in gen_conf.items() if k in LITELLM_ALLOWED_GEN_CONF_KEYS}
|
|
if deepseek_max_tokens is not None:
|
|
gen_conf["max_tokens"] = deepseek_max_tokens
|
|
return gen_conf
|
|
|
|
def _need_reasoning_content_back(self) -> bool:
|
|
return self.provider == SupportedLiteLLMProvider.DeepSeek
|
|
|
|
def _content_stream_sanitizer(self) -> "_StreamSanitizer | None":
|
|
"""A per-stream filter for providers whose control tokens leak into content."""
|
|
if self.provider == SupportedLiteLLMProvider.MiniMax:
|
|
return _StreamSanitizer(_MINIMAX_CONTROL_TOKEN_RE)
|
|
return None
|
|
|
|
def _sanitize_answer(self, text: str) -> str:
|
|
"""Strip provider control-token noise from a fully-assembled answer."""
|
|
if text and self.provider == SupportedLiteLLMProvider.MiniMax:
|
|
return _MINIMAX_CONTROL_TOKEN_RE.sub("", text)
|
|
return text
|
|
|
|
async def async_chat(self, system, history, gen_conf, **kwargs):
|
|
hist = list(history) if history else []
|
|
if system:
|
|
if not hist or hist[0].get("role") != "system":
|
|
hist.insert(0, {"role": "system", "content": system})
|
|
|
|
logging.info("[HISTORY]" + json.dumps(hist, ensure_ascii=False, indent=2))
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
_, kwargs = _apply_model_family_policies(
|
|
self.model_name,
|
|
backend="litellm",
|
|
provider=self.provider,
|
|
request_kwargs=kwargs,
|
|
)
|
|
|
|
completion_args = self._construct_completion_args(history=hist, stream=False, tools=False, **{**gen_conf, **kwargs})
|
|
|
|
for attempt in range(self.max_retries + 1):
|
|
try:
|
|
response = await litellm.acompletion(
|
|
**completion_args,
|
|
drop_params=True,
|
|
timeout=self.timeout,
|
|
)
|
|
|
|
# Capture the prompt/completion split for accurate per-call usage
|
|
# reporting (Langfuse + agent run aggregation).
|
|
self.last_usage = usage_from_response(response)
|
|
if not response.choices or not response.choices[0].message or not response.choices[0].message.content:
|
|
return "", 0
|
|
ans = response.choices[0].message.content.strip()
|
|
if response.choices[0].finish_reason == "length":
|
|
ans = self._length_stop(ans)
|
|
|
|
return ans, total_token_count_from_response(response)
|
|
except Exception as e:
|
|
e = await self._exceptions_async(e, attempt)
|
|
if e:
|
|
return e, 0
|
|
|
|
assert False, "Shouldn't be here."
|
|
|
|
async def async_chat_streamly(self, system, history, gen_conf, **kwargs):
|
|
if system and history and history[0].get("role") != "system":
|
|
history.insert(0, {"role": "system", "content": system})
|
|
logging.info("[HISTORY STREAMLY]" + json.dumps(history, ensure_ascii=False, indent=4))
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
reasoning_start = False
|
|
total_tokens = 0
|
|
# Reset so a stale split from a previous call can't leak into this one.
|
|
self.last_usage = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
|
|
|
completion_args = self._construct_completion_args(history=history, stream=True, tools=False, **gen_conf)
|
|
stop = kwargs.get("stop")
|
|
if stop:
|
|
completion_args["stop"] = stop
|
|
# Ask the provider to include authoritative usage in the final streaming chunk.
|
|
# drop_params=True ensures this is silently ignored by providers that don't support it.
|
|
completion_args.setdefault("stream_options", {})["include_usage"] = True
|
|
|
|
for attempt in range(self.max_retries + 1):
|
|
try:
|
|
stream = await litellm.acompletion(
|
|
**completion_args,
|
|
drop_params=True,
|
|
timeout=self.timeout,
|
|
)
|
|
|
|
async for resp in stream:
|
|
# Authoritative usage may arrive on a usage-only final chunk that
|
|
# carries no choices (OpenAI/OpenRouter with include_usage). Read it
|
|
# before the choices guard so the prompt/completion split is captured.
|
|
_usage = usage_from_response(resp)
|
|
if _usage["total_tokens"]:
|
|
total_tokens = _usage["total_tokens"]
|
|
self.last_usage = _usage
|
|
|
|
if not hasattr(resp, "choices") or not resp.choices:
|
|
continue
|
|
|
|
delta = resp.choices[0].delta
|
|
if not hasattr(delta, "content") or delta.content is None:
|
|
delta.content = ""
|
|
|
|
_reasoning = getattr(delta, "reasoning_content", None) or getattr(delta, "reasoning", None)
|
|
if kwargs.get("with_reasoning", True) and _reasoning:
|
|
ans = ""
|
|
if not reasoning_start:
|
|
reasoning_start = True
|
|
ans = "<think>"
|
|
ans += _reasoning + "</think>"
|
|
else:
|
|
reasoning_start = False
|
|
ans = delta.content
|
|
|
|
if not _usage["total_tokens"]:
|
|
# No authoritative usage yet: keep a running estimate as fallback.
|
|
total_tokens += num_tokens_from_string(delta.content)
|
|
|
|
finish_reason = resp.choices[0].finish_reason if hasattr(resp.choices[0], "finish_reason") else ""
|
|
if finish_reason == "length":
|
|
if is_chinese(ans):
|
|
ans += LENGTH_NOTIFICATION_CN
|
|
else:
|
|
ans += LENGTH_NOTIFICATION_EN
|
|
|
|
yield ans
|
|
yield total_tokens
|
|
return
|
|
except Exception as e:
|
|
e = await self._exceptions_async(e, attempt)
|
|
if e:
|
|
yield e
|
|
yield total_tokens
|
|
return
|
|
|
|
def _length_stop(self, ans):
|
|
if is_chinese([ans]):
|
|
return ans + LENGTH_NOTIFICATION_CN
|
|
return ans + LENGTH_NOTIFICATION_EN
|
|
|
|
@property
|
|
def _retryable_errors(self) -> set[str]:
|
|
return {
|
|
LLMErrorCode.ERROR_RATE_LIMIT,
|
|
LLMErrorCode.ERROR_SERVER,
|
|
}
|
|
|
|
def _should_retry(self, error_code: str) -> bool:
|
|
return error_code in self._retryable_errors
|
|
|
|
async def _exceptions_async(self, e, attempt):
|
|
logging.exception("LiteLLMBase async completion")
|
|
error_code = self._classify_error(e)
|
|
if attempt == self.max_retries:
|
|
error_code = LLMErrorCode.ERROR_MAX_RETRIES
|
|
|
|
if self._should_retry(error_code):
|
|
delay = self._get_delay()
|
|
logging.warning(f"Error: {error_code}. Retrying in {delay:.2f} seconds... (Attempt {attempt + 1}/{self.max_retries})")
|
|
await asyncio.sleep(delay)
|
|
return None
|
|
error_detail = str(e)
|
|
if self.provider == SupportedLiteLLMProvider.Nvidia and "function" in error_detail.lower() and "not found for account" in error_detail.lower():
|
|
model_name = self.model_name.removeprefix(self.prefix)
|
|
error_detail = f"NVIDIA hosted endpoint '{model_name}' is unavailable or deprecated; refresh the provider model list and select an active Free Endpoint. Original error: {error_detail}"
|
|
msg = f"{ERROR_PREFIX}: {error_code} - {error_detail}"
|
|
logging.error(f"async_chat_streamly giving up: {msg}")
|
|
return msg
|
|
|
|
def _verbose_tool_use(self, name, args, res):
|
|
return (
|
|
"<tool_call>"
|
|
+ json.dumps(
|
|
{"name": name, "args": args, "result": str(res) if isinstance(res, Exception) else res},
|
|
ensure_ascii=False,
|
|
indent=2,
|
|
)
|
|
+ "</tool_call>"
|
|
)
|
|
|
|
def _append_history(self, hist, tool_call, tool_res, reasoning_content=None):
|
|
assistant_msg = {
|
|
"role": "assistant",
|
|
"content": None,
|
|
"tool_calls": [
|
|
{
|
|
"index": getattr(tool_call, "index", None),
|
|
"id": tool_call.id,
|
|
"function": {
|
|
"name": tool_call.function.name,
|
|
"arguments": tool_call.function.arguments,
|
|
},
|
|
"type": "function",
|
|
},
|
|
],
|
|
}
|
|
if reasoning_content:
|
|
assistant_msg["reasoning_content"] = reasoning_content
|
|
hist.append(assistant_msg)
|
|
try:
|
|
if isinstance(tool_res, dict):
|
|
tool_res = json.dumps(tool_res, ensure_ascii=False)
|
|
finally:
|
|
hist.append({"role": "tool", "tool_call_id": tool_call.id, "content": str(tool_res)})
|
|
return hist
|
|
|
|
def _append_history_batch(self, hist, results, reasoning_content=None):
|
|
"""
|
|
Append a batch of tool calls to history following the OpenAI protocol:
|
|
one assistant message containing all tool_calls, followed by one tool message per call.
|
|
results: list of (tool_call, name, args, result, error)
|
|
"""
|
|
assistant_msg = {
|
|
"role": "assistant",
|
|
"content": None,
|
|
"tool_calls": [
|
|
{
|
|
"index": getattr(tc, "index", None),
|
|
"id": tc.id,
|
|
"function": {"name": tc.function.name, "arguments": tc.function.arguments},
|
|
"type": "function",
|
|
}
|
|
for tc, _, _, _, _ in results
|
|
],
|
|
}
|
|
if reasoning_content:
|
|
assistant_msg["reasoning_content"] = reasoning_content
|
|
hist.append(assistant_msg)
|
|
for tc, _, _, result, err in results:
|
|
if err:
|
|
content = str(err)
|
|
elif isinstance(result, dict):
|
|
content = json.dumps(result, ensure_ascii=False)
|
|
else:
|
|
content = str(result)
|
|
hist.append({"role": "tool", "tool_call_id": tc.id, "content": content.replace("</think>", "") + ("</think>" if content.find("<think>") >= 0 else "")})
|
|
return hist
|
|
|
|
def bind_tools(self, toolcall_session=None, tools=None):
|
|
"""Register tools the LLM can call.
|
|
|
|
Two calling styles are accepted:
|
|
|
|
* Legacy: ``bind_tools(toolcall_session, tools_schemas)`` where
|
|
``toolcall_session`` implements :class:`ToolCallSession` and
|
|
``tools_schemas`` is a pre-built list of OpenAI function-schema
|
|
dicts (used by the agent/dialog layer).
|
|
* Decorator: ``bind_tools(tools=[fn1, fn2, ...])`` where each ``fn``
|
|
is decorated with :func:`rag.llm.tool_decorator.tool`. The session
|
|
and schemas are derived from the callables automatically.
|
|
"""
|
|
if tools is None and isinstance(toolcall_session, list):
|
|
tools, toolcall_session = toolcall_session, None
|
|
|
|
if tools and toolcall_session is None and all(is_tool(t) for t in tools):
|
|
session = FunctionToolSession(tools)
|
|
self.is_tools = True
|
|
self.toolcall_session = session
|
|
self.tools = session.schemas
|
|
return
|
|
|
|
if not (toolcall_session and tools):
|
|
return
|
|
self.is_tools = True
|
|
self.toolcall_session = toolcall_session
|
|
self.tools = tools
|
|
|
|
async def async_chat_with_tools(self, system: str, history: list, gen_conf: dict | None = None):
|
|
gen_conf = dict(gen_conf or {})
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
if system and history and history[0].get("role") != "system":
|
|
history.insert(0, {"role": "system", "content": system})
|
|
|
|
ans = ""
|
|
tk_count = 0
|
|
# Aggregate prompt/completion/total across every tool-calling round so the
|
|
# whole multi-round exchange is reported once with a correct split.
|
|
agg_usage = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
|
|
|
def _add_usage(resp):
|
|
nonlocal tk_count
|
|
u = usage_from_response(resp)
|
|
agg_usage["prompt_tokens"] += u["prompt_tokens"]
|
|
agg_usage["completion_tokens"] += u["completion_tokens"]
|
|
agg_usage["total_tokens"] += u["total_tokens"] or total_token_count_from_response(resp)
|
|
tk_count = agg_usage["total_tokens"]
|
|
self.last_usage = dict(agg_usage)
|
|
|
|
hist = deepcopy(history)
|
|
for attempt in range(self.max_retries + 1):
|
|
history = deepcopy(hist)
|
|
try:
|
|
for _ in range(self.max_rounds + 1):
|
|
logging.info(f"HAS TOOL:{len(self.tools)}\n{history=}")
|
|
|
|
completion_args = self._construct_completion_args(history=history, stream=False, tools=True, **gen_conf)
|
|
response = await litellm.acompletion(
|
|
**completion_args,
|
|
drop_params=True,
|
|
timeout=self.timeout,
|
|
)
|
|
|
|
_add_usage(response)
|
|
|
|
if not hasattr(response, "choices") or not response.choices or not response.choices[0].message:
|
|
raise Exception(f"500 response structure error. Response: {response}")
|
|
|
|
message = response.choices[0].message
|
|
reasoning_content = None
|
|
if self._need_reasoning_content_back():
|
|
reasoning_content = getattr(message, "reasoning_content", None) or getattr(message, "reasoning", None)
|
|
|
|
if not hasattr(message, "tool_calls") or not message.tool_calls:
|
|
if reasoning_content:
|
|
ans += f"<think>{reasoning_content}</think>"
|
|
ans += message.content or ""
|
|
if response.choices[0].finish_reason == "length":
|
|
ans = self._length_stop(ans)
|
|
return self._sanitize_answer(ans), tk_count
|
|
|
|
async def _exec_tool(tc):
|
|
name = tc.function.name
|
|
try:
|
|
args = json_repair.loads(tc.function.arguments)
|
|
if not isinstance(args, dict):
|
|
raise TypeError(f"Tool arguments for {name} must be a JSON object, got {type(args).__name__}")
|
|
if hasattr(self.toolcall_session, "tool_call_async"):
|
|
result = await self.toolcall_session.tool_call_async(name, args)
|
|
else:
|
|
result = await thread_pool_exec(self.toolcall_session.tool_call, name, args)
|
|
return tc, name, args, result, None
|
|
except Exception as e:
|
|
logging.exception(f"Tool call failed: {tc}")
|
|
return tc, name, {}, None, e
|
|
|
|
logging.info(f"Response tool_calls={message.tool_calls}")
|
|
results = await asyncio.gather(*[_exec_tool(tc) for tc in message.tool_calls])
|
|
history = self._append_history_batch(
|
|
history,
|
|
results,
|
|
reasoning_content=reasoning_content if self._need_reasoning_content_back() else None,
|
|
)
|
|
for tc, name, args, result, err in results:
|
|
ans += self._verbose_tool_use(name, args, err if err else result)
|
|
|
|
logging.warning(f"Exceed max rounds: {self.max_rounds}")
|
|
history.append({"role": "user", "content": f"Exceed max rounds: {self.max_rounds}"})
|
|
|
|
response, token_count = await self.async_chat("", history, gen_conf)
|
|
ans += response
|
|
# self.async_chat set self.last_usage to its own call; fold it into the aggregate.
|
|
_fb = getattr(self, "last_usage", None) or {}
|
|
agg_usage["prompt_tokens"] += int(_fb.get("prompt_tokens", 0) or 0)
|
|
agg_usage["completion_tokens"] += int(_fb.get("completion_tokens", 0) or 0)
|
|
agg_usage["total_tokens"] += int(_fb.get("total_tokens", 0) or token_count)
|
|
tk_count = agg_usage["total_tokens"]
|
|
self.last_usage = dict(agg_usage)
|
|
return self._sanitize_answer(ans), tk_count
|
|
|
|
except Exception as e:
|
|
e = await self._exceptions_async(e, attempt)
|
|
if e:
|
|
return e, tk_count
|
|
|
|
assert False, "Shouldn't be here."
|
|
|
|
async def async_chat_streamly_with_tools(self, system: str, history: list, gen_conf: dict | None = None):
|
|
gen_conf = dict(gen_conf or {})
|
|
gen_conf = self._clean_conf(gen_conf)
|
|
tools = self.tools
|
|
if system and history and history[0].get("role") != "system":
|
|
history.insert(0, {"role": "system", "content": system})
|
|
|
|
total_tokens = 0
|
|
# Aggregate usage across every tool-calling round (each round is a separate
|
|
# provider request). Committing per round avoids the previous bug where a later
|
|
# round's total overwrote earlier rounds.
|
|
agg_usage = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
|
hist = deepcopy(history)
|
|
|
|
def _commit_round(round_usage, round_estimate):
|
|
nonlocal total_tokens
|
|
if round_usage and round_usage["total_tokens"]:
|
|
agg_usage["prompt_tokens"] += round_usage["prompt_tokens"]
|
|
agg_usage["completion_tokens"] += round_usage["completion_tokens"]
|
|
agg_usage["total_tokens"] += round_usage["total_tokens"]
|
|
else:
|
|
agg_usage["total_tokens"] += round_estimate
|
|
total_tokens = agg_usage["total_tokens"]
|
|
self.last_usage = dict(agg_usage)
|
|
|
|
for attempt in range(self.max_retries + 1):
|
|
history = deepcopy(hist)
|
|
try:
|
|
for _round in range(self.max_rounds + 1):
|
|
reasoning_start = False
|
|
reasoning_content = ""
|
|
logging.info(f"[Tool loop] Deciding what to do next (step {_round + 1}); available tools: {', '.join(t['function']['name'] for t in tools)}")
|
|
|
|
completion_args = self._construct_completion_args(history=history, stream=True, tools=True, **gen_conf)
|
|
# Request authoritative usage on the final streaming chunk.
|
|
completion_args.setdefault("stream_options", {})["include_usage"] = True
|
|
response = await litellm.acompletion(
|
|
**completion_args,
|
|
drop_params=True,
|
|
timeout=self.timeout,
|
|
)
|
|
|
|
final_tool_calls = {}
|
|
answer = ""
|
|
round_usage = None
|
|
round_estimate = 0
|
|
# Per-round filter for providers (MiniMax) whose control tokens
|
|
# leak into content split across deltas; None for others.
|
|
_sanitizer = self._content_stream_sanitizer()
|
|
|
|
async for resp in response:
|
|
# Usage-only final chunk may carry no choices — read it first.
|
|
_u = usage_from_response(resp)
|
|
if _u["total_tokens"]:
|
|
round_usage = _u
|
|
|
|
if not hasattr(resp, "choices") or not resp.choices:
|
|
continue
|
|
|
|
delta = resp.choices[0].delta
|
|
|
|
if hasattr(delta, "tool_calls") and delta.tool_calls:
|
|
for tool_call in delta.tool_calls:
|
|
index = tool_call.index
|
|
if index not in final_tool_calls:
|
|
if not tool_call.function.arguments:
|
|
tool_call.function.arguments = ""
|
|
final_tool_calls[index] = tool_call
|
|
else:
|
|
final_tool_calls[index].function.arguments += tool_call.function.arguments or ""
|
|
continue
|
|
|
|
if not hasattr(delta, "content") or delta.content is None:
|
|
delta.content = ""
|
|
|
|
_reasoning = getattr(delta, "reasoning_content", None) or getattr(delta, "reasoning", None)
|
|
if _reasoning:
|
|
if self._need_reasoning_content_back():
|
|
reasoning_content += _reasoning
|
|
ans = ""
|
|
if not reasoning_start:
|
|
reasoning_start = True
|
|
ans = "<think>"
|
|
ans += _reasoning + "</think>"
|
|
yield ans
|
|
else:
|
|
reasoning_start = False
|
|
answer += delta.content
|
|
if _sanitizer is not None:
|
|
emitted = _sanitizer.feed(delta.content)
|
|
if emitted:
|
|
yield emitted
|
|
else:
|
|
yield delta.content
|
|
|
|
if not _u["total_tokens"]:
|
|
round_estimate += num_tokens_from_string(delta.content)
|
|
|
|
finish_reason = getattr(resp.choices[0], "finish_reason", "")
|
|
if finish_reason == "length":
|
|
yield self._length_stop("")
|
|
|
|
# Flush any held-back (sanitized) answer content for this round.
|
|
if _sanitizer is not None:
|
|
tail = _sanitizer.flush()
|
|
if tail:
|
|
yield tail
|
|
|
|
# Commit this round's tokens to the running aggregate.
|
|
_commit_round(round_usage, round_estimate)
|
|
|
|
if answer and not final_tool_calls:
|
|
logging.info(f"[Tool loop] Answering directly at step {_round + 1} — no tool needed.")
|
|
yield total_tokens
|
|
return
|
|
|
|
async def _exec_tool(tc):
|
|
name = tc.function.name
|
|
try:
|
|
args = json_repair.loads(tc.function.arguments)
|
|
if not isinstance(args, dict):
|
|
raise TypeError(f"Tool arguments for {name} must be a JSON object, got {type(args).__name__}")
|
|
if hasattr(self.toolcall_session, "tool_call_async"):
|
|
result = await self.toolcall_session.tool_call_async(name, args)
|
|
else:
|
|
result = await thread_pool_exec(self.toolcall_session.tool_call, name, args)
|
|
return tc, name, args, result, None
|
|
except Exception as e:
|
|
logging.exception(f"Tool call failed: {tc}")
|
|
return tc, name, {}, None, e
|
|
|
|
tcs = list(final_tool_calls.values())
|
|
logging.info(f"[Tool loop] Step {_round + 1}: running {', '.join(tc.function.name for tc in tcs)}...")
|
|
for tc in tcs:
|
|
try:
|
|
args = json_repair.loads(tc.function.arguments)
|
|
except Exception:
|
|
args = {}
|
|
yield f"<think>Running the {tc.function.name} tool...</think>"
|
|
results = await asyncio.gather(*[_exec_tool(tc) for tc in tcs])
|
|
|
|
# Terminal-tool short-circuit: a terminal tool already
|
|
# produces the final answer, so stream its result and stop
|
|
# instead of feeding it back for another LLM round.
|
|
_terminal = getattr(self, "terminal_tools", None)
|
|
if _terminal:
|
|
for tc, name, args, result, err in results:
|
|
if name in _terminal and not err:
|
|
logging.info(f"[Tool loop] The {name} tool produced the final answer — done.")
|
|
out = result if isinstance(result, str) else json.dumps(result, ensure_ascii=False)
|
|
if out:
|
|
yield out
|
|
yield total_tokens
|
|
return
|
|
|
|
history = self._append_history_batch(
|
|
history,
|
|
results,
|
|
reasoning_content=reasoning_content if self._need_reasoning_content_back() else None,
|
|
)
|
|
for tc, name, args, result, err in results:
|
|
yield self._verbose_tool_use(name, args, err if err else result)
|
|
|
|
logging.warning(f"Exceed max rounds: {self.max_rounds}")
|
|
history.append({"role": "user", "content": f"Exceed max rounds: {self.max_rounds}"})
|
|
|
|
completion_args = self._construct_completion_args(history=history, stream=True, tools=True, **gen_conf)
|
|
completion_args.setdefault("stream_options", {})["include_usage"] = True
|
|
response = await litellm.acompletion(
|
|
**completion_args,
|
|
drop_params=True,
|
|
timeout=self.timeout,
|
|
)
|
|
|
|
fb_usage = None
|
|
fb_estimate = 0
|
|
async for resp in response:
|
|
_u = usage_from_response(resp)
|
|
if _u["total_tokens"]:
|
|
fb_usage = _u
|
|
if not hasattr(resp, "choices") or not resp.choices:
|
|
continue
|
|
delta = resp.choices[0].delta
|
|
if not hasattr(delta, "content") or delta.content is None:
|
|
continue
|
|
if not _u["total_tokens"]:
|
|
fb_estimate += num_tokens_from_string(delta.content)
|
|
yield delta.content
|
|
|
|
_commit_round(fb_usage, fb_estimate)
|
|
yield total_tokens
|
|
return
|
|
|
|
except Exception as e:
|
|
e = await self._exceptions_async(e, attempt)
|
|
if e:
|
|
yield e
|
|
yield total_tokens
|
|
return
|
|
|
|
assert False, "Shouldn't be here."
|
|
|
|
def _construct_completion_args(self, history, stream: bool, tools: bool, **kwargs):
|
|
completion_args = {
|
|
"model": self.model_name,
|
|
"messages": history,
|
|
"api_key": self.api_key,
|
|
**kwargs,
|
|
}
|
|
if self.provider == SupportedLiteLLMProvider.Nvidia:
|
|
completion_args["num_retries"] = 0
|
|
else:
|
|
completion_args.setdefault("num_retries", self.max_retries)
|
|
# Forward the originating session/user as the OpenAI-standard `user` field so
|
|
# providers (OpenAI, OpenRouter, ...) receive it in the request body and
|
|
# upstream activity can be correlated back to the session. An explicit
|
|
# caller-supplied `user` (including an empty string to suppress it) wins, so
|
|
# check key presence rather than truthiness.
|
|
if "user" not in completion_args:
|
|
request_user = current_llm_user()
|
|
if request_user:
|
|
completion_args["user"] = request_user
|
|
if stream:
|
|
completion_args.update(
|
|
{
|
|
"stream": stream,
|
|
}
|
|
)
|
|
if tools and self.tools:
|
|
completion_args.update(
|
|
{
|
|
"tools": self.tools,
|
|
"tool_choice": "auto",
|
|
}
|
|
)
|
|
if self.provider in FACTORY_DEFAULT_BASE_URL:
|
|
completion_args.update({"api_base": self.base_url})
|
|
elif self.provider == SupportedLiteLLMProvider.Bedrock:
|
|
import boto3
|
|
|
|
completion_args.pop("api_key", None)
|
|
completion_args.pop("api_base", None)
|
|
|
|
bedrock_key = json.loads(self.api_key)
|
|
mode = bedrock_key.get("auth_mode")
|
|
if not mode:
|
|
logging.error("Bedrock auth_mode is not provided in the key")
|
|
raise ValueError("Bedrock auth_mode must be provided in the key")
|
|
|
|
bedrock_region = bedrock_key.get("bedrock_region")
|
|
|
|
if mode == "access_key_secret":
|
|
completion_args.update({"aws_region_name": bedrock_region})
|
|
completion_args.update({"aws_access_key_id": bedrock_key.get("bedrock_ak")})
|
|
completion_args.update({"aws_secret_access_key": bedrock_key.get("bedrock_sk")})
|
|
elif mode == "iam_role":
|
|
aws_role_arn = bedrock_key.get("aws_role_arn")
|
|
sts_client = boto3.client("sts", region_name=bedrock_region)
|
|
resp = sts_client.assume_role(RoleArn=aws_role_arn, RoleSessionName="BedrockSession")
|
|
creds = resp["Credentials"]
|
|
completion_args.update({"aws_region_name": bedrock_region})
|
|
completion_args.update({"aws_access_key_id": creds["AccessKeyId"]})
|
|
completion_args.update({"aws_secret_access_key": creds["SecretAccessKey"]})
|
|
completion_args.update({"aws_session_token": creds["SessionToken"]})
|
|
else: # assume_role - use default credential chain (IRSA, instance profile, etc.)
|
|
completion_args.update({"aws_region_name": bedrock_region})
|
|
|
|
elif self.provider == SupportedLiteLLMProvider.OpenRouter:
|
|
if self.provider_order:
|
|
|
|
def _to_order_list(x):
|
|
if x is None:
|
|
return []
|
|
if isinstance(x, str):
|
|
return [s.strip() for s in x.split(",") if s.strip()]
|
|
if isinstance(x, (list, tuple)):
|
|
return [str(s).strip() for s in x if str(s).strip()]
|
|
return []
|
|
|
|
extra_body = {}
|
|
provider_cfg = {}
|
|
provider_order = _to_order_list(self.provider_order)
|
|
provider_cfg["order"] = provider_order
|
|
provider_cfg["allow_fallbacks"] = False
|
|
extra_body["provider"] = provider_cfg
|
|
completion_args.update({"extra_body": extra_body})
|
|
elif self.provider == SupportedLiteLLMProvider.GPUStack:
|
|
completion_args.update(
|
|
{
|
|
"api_base": urljoin(self.base_url, "v1"),
|
|
}
|
|
)
|
|
elif self.provider == SupportedLiteLLMProvider.Azure_OpenAI:
|
|
completion_args.pop("api_key", None)
|
|
completion_args.pop("api_base", None)
|
|
completion_args.update(
|
|
{
|
|
"api_key": self.api_key,
|
|
"api_base": self.base_url,
|
|
"api_version": self.api_version,
|
|
}
|
|
)
|
|
|
|
# Ollama deployments commonly sit behind a reverse proxy that enforces
|
|
# Bearer auth. Ensure the Authorization header is set when an API key
|
|
# is provided, while respecting any user-supplied headers. #11350
|
|
extra_headers = deepcopy(completion_args.get("extra_headers") or {})
|
|
if self.provider == SupportedLiteLLMProvider.Ollama and self.api_key and "Authorization" not in extra_headers:
|
|
extra_headers["Authorization"] = f"Bearer {self.api_key}"
|
|
# MiniMax requires GroupId as a query parameter for API authentication
|
|
if self.provider == SupportedLiteLLMProvider.MiniMax and hasattr(self, "group_id") and self.group_id:
|
|
api_base = completion_args.get("api_base", self.base_url)
|
|
separator = "&" if "?" in api_base else "?"
|
|
completion_args["api_base"] = f"{api_base}{separator}GroupId={self.group_id}"
|
|
_move_litellm_provider_body_fields(self.provider, completion_args)
|
|
if extra_headers:
|
|
completion_args["extra_headers"] = extra_headers
|
|
return completion_args
|
|
|
|
|
|
class RAGconChat(Base):
|
|
"""
|
|
RAGcon Chat Provider - routes through LiteLLM proxy
|
|
|
|
All model types are handled through a unified LiteLLM endpoint.
|
|
Default Base URL: https://connect.ragcon.com/v1
|
|
"""
|
|
|
|
_FACTORY_NAME = "RAGcon"
|
|
|
|
def __init__(self, key, model_name, base_url=None, **kwargs):
|
|
if not base_url:
|
|
base_url = "https://connect.ragcon.com/v1"
|
|
|
|
super().__init__(key, model_name, base_url, **kwargs)
|
|
|
|
|
|
class NewAPIChat(Base):
|
|
_FACTORY_NAME = "New API"
|
|
|
|
def __init__(self, key, model_name, base_url, **kwargs):
|
|
if not base_url:
|
|
raise ValueError("url cannot be None")
|
|
model_name = model_name.split("___")[0]
|
|
super().__init__(key, model_name, base_url, **kwargs)
|