refactor: address code review feedback on MCP tool calling
This commit is contained in:
parent
76c6086db3
commit
51188fe3bd
@ -371,7 +371,7 @@ async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) ->
|
|||||||
if tool_client is not None:
|
if tool_client is not None:
|
||||||
reply = await llm.chat_with_tools(text, tool_client, history=call_history)
|
reply = await llm.chat_with_tools(text, tool_client, history=call_history)
|
||||||
else:
|
else:
|
||||||
reply = await llm.chat(text, history=call_history) # type: ignore[arg-type]
|
reply = await llm.chat(text, history=call_history)
|
||||||
history.append({"role": "user", "content": text})
|
history.append({"role": "user", "content": text})
|
||||||
history.append({"role": "assistant", "content": reply})
|
history.append({"role": "assistant", "content": reply})
|
||||||
else:
|
else:
|
||||||
@ -381,7 +381,7 @@ async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) ->
|
|||||||
if tool_client is not None:
|
if tool_client is not None:
|
||||||
reply = await llm.chat_with_tools(text, tool_client, history=call_history)
|
reply = await llm.chat_with_tools(text, tool_client, history=call_history)
|
||||||
else:
|
else:
|
||||||
reply = await llm.chat(text, history=call_history) # type: ignore[arg-type]
|
reply = await llm.chat(text, history=call_history)
|
||||||
user_history.append({"role": "user", "content": text})
|
user_history.append({"role": "user", "content": text})
|
||||||
user_history.append({"role": "assistant", "content": reply})
|
user_history.append({"role": "assistant", "content": reply})
|
||||||
if len(user_history) > _MAX_HISTORY * 2:
|
if len(user_history) > _MAX_HISTORY * 2:
|
||||||
|
|||||||
@ -29,12 +29,12 @@ class LLMClient:
|
|||||||
self,
|
self,
|
||||||
user_message: str,
|
user_message: str,
|
||||||
*,
|
*,
|
||||||
history: list[dict[str, str]] | None = None,
|
history: list[dict[str, Any]] | None = None,
|
||||||
system_prompt: str | None = None,
|
system_prompt: str | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Send a user message (with optional history) and return the assistant reply."""
|
"""Send a user message (with optional history) and return the assistant reply."""
|
||||||
system = system_prompt or self._settings.openai_system_prompt
|
system = system_prompt or self._settings.openai_system_prompt
|
||||||
messages: list[dict[str, str]] = [{"role": "system", "content": system}]
|
messages: list[dict[str, Any]] = [{"role": "system", "content": system}]
|
||||||
if history:
|
if history:
|
||||||
messages.extend(history)
|
messages.extend(history)
|
||||||
messages.append({"role": "user", "content": user_message})
|
messages.append({"role": "user", "content": user_message})
|
||||||
@ -130,8 +130,13 @@ class LLMClient:
|
|||||||
|
|
||||||
# Execute each tool and append results
|
# Execute each tool and append results
|
||||||
for tc in msg.tool_calls:
|
for tc in msg.tool_calls:
|
||||||
|
try:
|
||||||
try:
|
try:
|
||||||
args = json.loads(tc.function.arguments)
|
args = json.loads(tc.function.arguments)
|
||||||
|
except json.JSONDecodeError as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"Invalid JSON in tool arguments for '{tc.function.name}': {exc}"
|
||||||
|
) from exc
|
||||||
result = await tool_client.call(tc.function.name, args)
|
result = await tool_client.call(tc.function.name, args)
|
||||||
logger.info("Tool %s → %d chars", tc.function.name, len(result))
|
logger.info("Tool %s → %d chars", tc.function.name, len(result))
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
@ -153,12 +158,12 @@ class LLMClient:
|
|||||||
self,
|
self,
|
||||||
user_message: str,
|
user_message: str,
|
||||||
*,
|
*,
|
||||||
history: list[dict[str, str]] | None = None,
|
history: list[dict[str, Any]] | None = None,
|
||||||
system_prompt: str | None = None,
|
system_prompt: str | None = None,
|
||||||
) -> AsyncIterator[str]:
|
) -> AsyncIterator[str]:
|
||||||
"""Stream the assistant reply token by token."""
|
"""Stream the assistant reply token by token."""
|
||||||
system = system_prompt or self._settings.openai_system_prompt
|
system = system_prompt or self._settings.openai_system_prompt
|
||||||
messages: list[dict[str, str]] = [{"role": "system", "content": system}]
|
messages: list[dict[str, Any]] = [{"role": "system", "content": system}]
|
||||||
if history:
|
if history:
|
||||||
messages.extend(history)
|
messages.extend(history)
|
||||||
messages.append({"role": "user", "content": user_message})
|
messages.append({"role": "user", "content": user_message})
|
||||||
|
|||||||
@ -40,7 +40,7 @@ class ToolClient:
|
|||||||
def _headers(self) -> dict[str, str]:
|
def _headers(self) -> dict[str, str]:
|
||||||
headers: dict[str, str] = {"Accept": "application/json"}
|
headers: dict[str, str] = {"Accept": "application/json"}
|
||||||
if self._api_key:
|
if self._api_key:
|
||||||
headers["Authorization"] = "Bearer " + self._api_key
|
headers["Authorization"] = "Bearer " + self._api_key.strip()
|
||||||
return headers
|
return headers
|
||||||
|
|
||||||
async def get_spec(self) -> dict[str, Any]:
|
async def get_spec(self) -> dict[str, Any]:
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user