Merge pull request #51 from HKUDS/review/xzq-batch-20260403

Review/xzq batch 20260403
This commit is contained in:
Dennis-yxchen 2026-04-03 23:51:09 +08:00 committed by GitHub
commit 456184f6e1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 76 additions and 25 deletions

View file

@ -337,13 +337,13 @@ class GroundingClient:
def get_session_info(self, name: str) -> SessionInfo:
"""Get session monitoring info"""
if name not in self._session_info:
raise ErrorCode.SESSION_NOT_FOUND(name)
raise GroundingError(f"Session not found: {name}", code=ErrorCode.SESSION_NOT_FOUND)
return self._session_info[name]
def get_session(self, name: str) -> BaseSession:
"""Get session"""
if name not in self._sessions:
raise ErrorCode.SESSION_NOT_FOUND(name)
raise GroundingError(f"Session not found: {name}", code=ErrorCode.SESSION_NOT_FOUND)
return self._sessions[name]
@ -479,7 +479,7 @@ class GroundingClient:
# Session-level
if session_name:
if session_name not in self._sessions:
raise ErrorCode.SESSION_NOT_FOUND(session_name)
raise GroundingError(f"Session not found: {session_name}", code=ErrorCode.SESSION_NOT_FOUND)
backend_type = self._session_info[session_name].backend_type
return await self._fetch_tools(
backend_type,
@ -531,7 +531,7 @@ class GroundingClient:
use_cache: bool = False
) -> list[BaseTool]:
if session_name not in self._session_info:
raise ErrorCode.SESSION_NOT_FOUND(session_name)
raise GroundingError(f"Session not found: {session_name}", code=ErrorCode.SESSION_NOT_FOUND)
backend = self._session_info[session_name].backend_type
return await self.list_tools(backend, session_name, use_cache)
@ -838,7 +838,7 @@ class GroundingClient:
runtime_backend = backend
else:
if runtime_session not in self._session_info:
raise ErrorCode.SESSION_NOT_FOUND(runtime_session)
raise GroundingError(f"Session not found: {runtime_session}", code=ErrorCode.SESSION_NOT_FOUND)
runtime_backend = self._session_info[
runtime_session
].backend_type

View file

@ -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
@ -754,15 +780,24 @@ class LLMClient:
except:
pass
if tool_name not in tool_map:
result = ToolResult(
status=ToolStatus.ERROR,
error=f"Tool '{tool_name}' not found"
)
if tool_obj is None:
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(
tool=tool_map[tool_name],
tool=tool_obj,
openai_tool_call={
"id": tool_call.id,
"type": "function",
@ -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

View file

@ -22,7 +22,6 @@ import json
import logging
import os
import sys
import traceback
from pathlib import Path
from typing import Any, Dict, List, Optional
@ -597,7 +596,7 @@ async def execute_task(
except Exception as e:
logger.error(f"execute_task failed: {e}", exc_info=True)
return _json_error(e, status="error", traceback=traceback.format_exc(limit=5))
return _json_error(e, status="error")
@mcp.tool()
@ -818,7 +817,7 @@ async def fix_skill(
except Exception as e:
logger.error(f"fix_skill failed: {e}", exc_info=True)
return _json_error(e, status="error", traceback=traceback.format_exc(limit=5))
return _json_error(e, status="error")
@mcp.tool()
@ -889,7 +888,7 @@ async def upload_skill(
except Exception as e:
logger.error(f"upload_skill failed: {e}", exc_info=True)
return _json_error(e, status="error", traceback=traceback.format_exc(limit=5))
return _json_error(e, status="error")
def run_mcp_server() -> None:
"""Console-script entry point for ``openspace-mcp``."""

View file

@ -233,6 +233,23 @@ class Logger:
cls._configured = True
@classmethod
def set_level(cls, level: str) -> None:
"""Set log level by name (e.g. ``"DEBUG"``, ``"INFO"``, ``"WARNING"``)."""
resolved = getattr(logging, level.upper(), None)
if resolved is None or not isinstance(resolved, int):
raise ValueError(f"Unknown log level: {level!r}")
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:
"""Dynamically switch debug level: 0 = WARNING, 1 = INFO, 2 = DEBUG."""
@ -309,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")