mirror of
https://github.com/HKUDS/OpenSpace.git
synced 2026-09-11 22:51:05 +00:00
fix: make tool fallback conservative
Co-authored-by: xzq.xu <zhiqiang.xu@nodeskai.com>
This commit is contained in:
parent
af1eb5bbe6
commit
34d82b735e
2 changed files with 59 additions and 15 deletions
|
|
@ -190,6 +190,37 @@ def _infer_backend_from_tool_name(tool_name: str) -> Optional[str]:
|
|||
return None
|
||||
|
||||
|
||||
def _resolve_tool_call_target(
|
||||
tool_name: str,
|
||||
tool_map: Dict[str, BaseTool],
|
||||
) -> tuple[Optional[BaseTool], List[str]]:
|
||||
"""Resolve a returned tool name to a concrete tool object.
|
||||
|
||||
The LLM is expected to return the deduped tool key from ``tool_map``.
|
||||
Some providers occasionally return the short schema name instead. In that
|
||||
case we only recover when exactly one tool shares that schema name; if
|
||||
multiple tools match, the call is ambiguous and should not be executed.
|
||||
"""
|
||||
tool_obj = tool_map.get(tool_name)
|
||||
if tool_obj is not None or not tool_name:
|
||||
return tool_obj, []
|
||||
|
||||
fallback_matches = [
|
||||
(llm_name, tool)
|
||||
for llm_name, tool in tool_map.items()
|
||||
if getattr(getattr(tool, "schema", None), "name", None) == tool_name
|
||||
]
|
||||
if len(fallback_matches) == 1:
|
||||
resolved_name, resolved_tool = fallback_matches[0]
|
||||
logger.info(
|
||||
f"[TOOL_FALLBACK] Resolved short tool name '{tool_name}' to '{resolved_name}'"
|
||||
)
|
||||
return resolved_tool, []
|
||||
if len(fallback_matches) > 1:
|
||||
return None, [llm_name for llm_name, _tool in fallback_matches]
|
||||
return None, []
|
||||
|
||||
|
||||
DEFAULT_SUMMARIZE_THRESHOLD_CHARS = 200000 # ~50K tokens, lowered from 400K to prevent context overflow
|
||||
MAX_TOOL_RESULT_CHARS = 200000 # Fallback truncation limit when summarization fails (~50K tokens)
|
||||
|
||||
|
|
@ -705,14 +736,9 @@ class LLMClient:
|
|||
for tool_call in tool_calls:
|
||||
tool_name = tool_call.function.name
|
||||
|
||||
# Resolve tool instance: key might differ from model response (e.g. API returns
|
||||
# "read_file" while we stored "server__read_file" for dedup), so fallback by schema.name
|
||||
tool_obj = tool_map.get(tool_name)
|
||||
if tool_obj is None and tool_name:
|
||||
for _k, _t in tool_map.items():
|
||||
if getattr(getattr(_t, "schema", None), "name", None) == tool_name:
|
||||
tool_obj = _t
|
||||
break
|
||||
# Resolve tool instance: some providers return the short schema
|
||||
# name instead of the deduped LLM-visible tool key.
|
||||
tool_obj, ambiguous_tool_names = _resolve_tool_call_target(tool_name, tool_map)
|
||||
|
||||
backend = None
|
||||
server_name = None
|
||||
|
|
@ -755,10 +781,19 @@ class LLMClient:
|
|||
pass
|
||||
|
||||
if tool_obj is None:
|
||||
result = ToolResult(
|
||||
status=ToolStatus.ERROR,
|
||||
error=f"Tool '{tool_name}' not found"
|
||||
)
|
||||
if ambiguous_tool_names:
|
||||
result = ToolResult(
|
||||
status=ToolStatus.ERROR,
|
||||
error=(
|
||||
f"Tool '{tool_name}' is ambiguous; matches: "
|
||||
f"{', '.join(ambiguous_tool_names)}"
|
||||
)
|
||||
)
|
||||
else:
|
||||
result = ToolResult(
|
||||
status=ToolStatus.ERROR,
|
||||
error=f"Tool '{tool_name}' not found"
|
||||
)
|
||||
else:
|
||||
try:
|
||||
result = await _execute_tool_call(
|
||||
|
|
@ -871,4 +906,4 @@ class LLMClient:
|
|||
role = msg.get("role", "unknown").upper()
|
||||
content = msg.get("content", "")
|
||||
formatted += f"[{role}]\n{content}\n\n"
|
||||
return formatted
|
||||
return formatted
|
||||
|
|
|
|||
|
|
@ -239,7 +239,16 @@ class Logger:
|
|||
resolved = getattr(logging, level.upper(), None)
|
||||
if resolved is None or not isinstance(resolved, int):
|
||||
raise ValueError(f"Unknown log level: {level!r}")
|
||||
cls.configure(level=resolved, force=True)
|
||||
if not cls._configured:
|
||||
cls.configure(level=resolved, attach_to_root=True)
|
||||
return
|
||||
|
||||
root_logger = logging.getLogger()
|
||||
root_logger.setLevel(resolved)
|
||||
for handler in root_logger.handlers:
|
||||
handler.setLevel(resolved)
|
||||
|
||||
cls._update_level(resolved)
|
||||
|
||||
@classmethod
|
||||
def set_debug(cls, debug_level: int = 2) -> None:
|
||||
|
|
@ -317,4 +326,4 @@ Logger.configure(attach_to_root=True)
|
|||
|
||||
# Get openspace logger for internal logging
|
||||
logger = Logger.get_logger()
|
||||
logger.debug("OpenSpace logging initialized")
|
||||
logger.debug("OpenSpace logging initialized")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue