Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
92 changes: 90 additions & 2 deletions src/vtk_prompt/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,40 @@
logger = get_logger(__name__)


def _parse_text_tool_calls(content: str | None, tool_names: set[str]) -> list[dict] | None:
"""Extract tool calls a backend emitted as text instead of structured tool_calls.

Some local OpenAI-compatible backends (e.g. Ollama with a quantized model) return
a tool call as plain content rather than populating ``tool_calls``. Handle both a
``<tool_call>{...}</tool_call>`` block and a bare JSON object with ``name`` and
``arguments``. Only objects whose name matches a known tool are treated as calls,
so normal tagged answers are never misread. Returns a list of {name, arguments}.
"""
if not content:
return None
candidates = re.findall(r"<tool_call>\s*(\{.*?\})\s*</tool_call>", content, re.DOTALL)
if not candidates:
stripped = content.strip()
if stripped.startswith("{") and stripped.endswith("}"):
candidates = [stripped]
calls: list[dict] = []
for raw in candidates:
try:
obj = json.loads(raw)
except (ValueError, TypeError):
continue
name = obj.get("name")
args = obj.get("arguments", {})
if isinstance(args, str):
try:
args = json.loads(args)
except (ValueError, TypeError):
args = {}
if name in tool_names and isinstance(args, dict):
calls.append({"name": name, "arguments": args})
return calls or None


@dataclass
class VTKPromptClient:
"""OpenAI client for VTK code generation."""
Expand Down Expand Up @@ -296,6 +330,8 @@ def query(
custom_prompt: dict | None = None,
ui_mode: bool = False,
execution_error: str | None = None,
log_tool_calls: bool = False,
agentic_retrieval: bool = False,
) -> tuple[str, str, Any] | tuple[str, str, Any, list[str]] | str:
"""Generate VTK code using vtk-mcp tools when available.

Expand Down Expand Up @@ -359,7 +395,8 @@ def query(
else:
# Normal path: build context and prompt
context_snippets = None
if mcp_client:
# Agentic mode: skip pre-injected context so the model must use tools.
if mcp_client and not agentic_retrieval:
mcp_context = mcp_client.get_enriched_context(message, top_k=top_k)
if mcp_context:
context_snippets = mcp_context
Expand Down Expand Up @@ -417,6 +454,11 @@ def query(

# Fetch vtk-mcp tools for LLM tool calling
tools = mcp_client.list_tools() if mcp_client else []
tool_names: set[str] = {
str(name)
for t in tools
if (name := (t.get("function") or {}).get("name")) is not None
}

# Retry loop for AST validation
for attempt in range(retry_attempts):
Expand Down Expand Up @@ -468,18 +510,64 @@ def query(
except Exception:
args = {}
result = mcp_client.call_tool(tc.function.name, args) # type: ignore
logger.debug("Tool %s -> %s...", tc.function.name, result[:80])
if log_tool_calls:
logger.info(
"vtk-mcp: %s(%s) -> %s",
tc.function.name,
tc.function.arguments,
result[:120],
)
else:
logger.debug("Tool %s -> %s...", tc.function.name, result[:80])
self.conversation.append(
{"role": "tool", "tool_call_id": tc.id, "content": result}
)
continue # let LLM decide what to do next

# Fallback: the backend returned a tool call as plain text rather than
# in tool_calls (common with quantized local models). Run it anyway.
text_calls = _parse_text_tool_calls(choice.message.content, tool_names)
if tools and text_calls:
text_msg: dict = {"role": "assistant", "content": ""}
text_msg["tool_calls"] = [
{
"id": f"call_{i}",
"type": "function",
"function": {
"name": c["name"],
"arguments": json.dumps(c["arguments"]),
},
}
for i, c in enumerate(text_calls)
]
self.conversation.append(text_msg)
for i, c in enumerate(text_calls):
result = mcp_client.call_tool(c["name"], c["arguments"]) # type: ignore
if log_tool_calls:
logger.info(
"vtk-mcp (text): %s(%s) -> %s",
c["name"],
json.dumps(c["arguments"]),
result[:120],
)
else:
logger.debug("Tool %s (text) -> %s...", c["name"], result[:80])
self.conversation.append(
{"role": "tool", "tool_call_id": f"call_{i}", "content": result}
)
continue

# LLM generated a response (not a tool call)
content = choice.message.content or "No content in response"
break

if content is None:
# Tool loop exhausted without a text response
if log_tool_calls:
logger.info(
"vtk-mcp: tool loop hit the %d-round cap without a final response",
MAX_TOOL_ROUNDS,
)
if attempt == retry_attempts - 1:
return ("No response generated", "", getattr(response, "usage", None) or {})
continue
Expand Down
4 changes: 4 additions & 0 deletions src/vtk_prompt/controllers/configuration.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,8 @@ def save_config(app: Any) -> str:
mcp_url = getattr(app.state, "mcp_url", "").strip()
data_root = getattr(app.state, "data_root", "").strip()
top_k = int(getattr(app.state, "top_k", 5))
log_tool_calls = bool(getattr(app.state, "log_tool_calls", False))
agentic_retrieval = bool(getattr(app.state, "agentic_retrieval", False))
base_url = getattr(app.state, "local_base_url", "").strip() if not use_cloud else ""

content = {
Expand All @@ -56,6 +58,8 @@ def save_config(app: Any) -> str:
"mcp_url": mcp_url,
"data_root": data_root,
"top_k": top_k,
"log_tool_calls": log_tool_calls,
"agentic_retrieval": agentic_retrieval,
"retries": retries,
"modelParameters": {
"temperature": temperature,
Expand Down
4 changes: 4 additions & 0 deletions src/vtk_prompt/controllers/generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,6 +103,8 @@ async def generate_and_execute_code(app: Any) -> None:
temperature=float(app.state.temperature),
top_k=int(app.state.top_k),
retry_attempts=int(app.state.retry_attempts),
log_tool_calls=bool(app.state.log_tool_calls),
agentic_retrieval=bool(app.state.agentic_retrieval),
provider=app.state.provider,
custom_prompt=app.custom_prompt_data,
ui_mode=True, # This tells the client to use UI-specific components
Expand Down Expand Up @@ -176,6 +178,8 @@ async def generate_and_execute_code(app: Any) -> None:
temperature=float(app.state.temperature),
top_k=int(app.state.top_k),
retry_attempts=1,
log_tool_calls=bool(app.state.log_tool_calls),
agentic_retrieval=bool(app.state.agentic_retrieval),
provider=app.state.provider,
custom_prompt=app.custom_prompt_data,
ui_mode=True,
Expand Down
2 changes: 2 additions & 0 deletions src/vtk_prompt/state/initializer.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,8 @@ def initialize_state(app: Any) -> None:
app.state.code_history_pos = -1
app.state.is_loading = False
app.state.mcp_url = ""
app.state.log_tool_calls = False # log vtk-mcp tool calls to the console
app.state.agentic_retrieval = False # skip pre-injected context; use tools
app.state.error_message = ""
app.state.input_tokens = 0
app.state.output_tokens = 0
Expand Down
17 changes: 17 additions & 0 deletions src/vtk_prompt/ui/layout/settings_dialog.py
Original file line number Diff line number Diff line change
Expand Up @@ -190,6 +190,23 @@ def _advanced_tab() -> None:
hint="Context snippets retrieved per request",
persistent_hint=True,
)
vuetify.VCheckbox(
label="Log tool calls to the server console",
v_model=("log_tool_calls", False),
density="compact",
color="primary",
disabled=("!mcp_url",),
hide_details=True,
classes="mt-2",
)
vuetify.VCheckbox(
label="Agentic retrieval (use tools instead of pre-injected context)",
v_model=("agentic_retrieval", False),
density="compact",
color="primary",
disabled=("!mcp_url",),
hide_details=True,
)

vuetify.VDivider(classes="my-5")
_section("Generation")
Expand Down
8 changes: 8 additions & 0 deletions src/vtk_prompt/utils/prompt_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -120,6 +120,14 @@ def _process_rag_and_generation_settings(app: Any) -> None:
_mcp = app.custom_prompt_data.get("mcp_url")
if isinstance(_mcp, str):
app.state.mcp_url = _mcp.strip()
if "log_tool_calls" in app.custom_prompt_data:
_ltc = app.custom_prompt_data.get("log_tool_calls")
if isinstance(_ltc, bool):
app.state.log_tool_calls = _ltc
if "agentic_retrieval" in app.custom_prompt_data:
_ag = app.custom_prompt_data.get("agentic_retrieval")
if isinstance(_ag, bool):
app.state.agentic_retrieval = _ag
if "base_url" in app.custom_prompt_data:
_base = app.custom_prompt_data.get("base_url")
if isinstance(_base, str) and _base.strip():
Expand Down