diff --git a/frontend/src/components/Chat/InputArea.tsx b/frontend/src/components/Chat/InputArea.tsx index 7c7f5978..20cf7301 100644 --- a/frontend/src/components/Chat/InputArea.tsx +++ b/frontend/src/components/Chat/InputArea.tsx @@ -5,6 +5,7 @@ import { useAppStore, generateId } from '../../lib/store'; import { streamChat, streamResearch } from '../../lib/sse'; import { fetchSavings, getBase } from '../../lib/api'; import { listConnectors, getSyncStatus } from '../../lib/connectors-api'; +import { serializeToolCallArguments } from '../../lib/tool-call'; import { MicButton } from './MicButton'; import { useSpeech } from '../../hooks/useSpeech'; import type { @@ -389,7 +390,7 @@ export function InputArea() { const tc: ToolCallInfo = { id: generateId(), tool: data.tool, - arguments: data.arguments || '', + arguments: serializeToolCallArguments(data.arguments), status: 'running', }; toolCalls.push(tc); @@ -400,7 +401,7 @@ export function InputArea() { updateLastAssistant(convId, accumulatedContent, [...toolCalls]); useAppStore.getState().addLogEntry({ timestamp: Date.now(), level: 'info', category: 'tool', - message: `Calling ${data.tool}(${data.arguments || ''})`, + message: `Calling ${data.tool}(${serializeToolCallArguments(data.arguments)})`, }); } catch {} } else if (eventName === 'tool_call_end') { diff --git a/frontend/src/components/Chat/ToolCallCard.tsx b/frontend/src/components/Chat/ToolCallCard.tsx index eb3d6aa3..03419672 100644 --- a/frontend/src/components/Chat/ToolCallCard.tsx +++ b/frontend/src/components/Chat/ToolCallCard.tsx @@ -1,6 +1,7 @@ import { useState } from 'react'; import { ChevronDown, ChevronRight, Loader2, CheckCircle2, XCircle } from 'lucide-react'; import type { ToolCallInfo } from '../../types'; +import { serializeToolCallArguments } from '../../lib/tool-call'; interface Props { toolCall: ToolCallInfo; @@ -35,7 +36,10 @@ export function ToolCallCard({ toolCall }: Props) { const [expanded, setExpanded] = useState(false); const config = statusConfig[toolCall.status]; const StatusIcon = config.icon; - const preview = previewArgs(toolCall.arguments); + // Persisted conversations may contain the pre-fix object payload despite + // the TypeScript contract, so normalize again at the final render boundary. + const argumentsText = serializeToolCallArguments(toolCall.arguments); + const preview = previewArgs(argumentsText); return (
- {toolCall.arguments && ( + {argumentsText && (
- {formatJson(toolCall.arguments)} + {formatJson(argumentsText)}
)} diff --git a/frontend/src/lib/api.ts b/frontend/src/lib/api.ts index e635f771..56ff88a8 100644 --- a/frontend/src/lib/api.ts +++ b/frontend/src/lib/api.ts @@ -1,5 +1,6 @@ import type { ModelInfo, SavingsData, ServerInfo } from '../types'; import { SUPABASE_ANON_KEY, SUPABASE_URL } from './supabase'; +import { serializeToolCallArguments } from './tool-call'; // --------------------------------------------------------------------------- // Supabase config @@ -741,7 +742,7 @@ export async function sendAgentMessage( const parsed = JSON.parse(data); callbacks?.onToolCallStart?.({ tool: parsed.tool, - arguments: parsed.arguments ?? '', + arguments: serializeToolCallArguments(parsed.arguments), }); } catch { /* skip */ diff --git a/frontend/src/lib/store.tool-call-repair.test.ts b/frontend/src/lib/store.tool-call-repair.test.ts new file mode 100644 index 00000000..a88fedf0 --- /dev/null +++ b/frontend/src/lib/store.tool-call-repair.test.ts @@ -0,0 +1,122 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; + +const CONVERSATIONS_KEY = 'openjarvis-conversations'; + +class MemoryStorage { + private store = new Map(); + + getItem(key: string): string | null { + return this.store.get(key) ?? null; + } + + setItem(key: string, value: string): void { + this.store.set(key, String(value)); + } + + removeItem(key: string): void { + this.store.delete(key); + } +} + +beforeEach(() => { + vi.resetModules(); + (globalThis as unknown as { localStorage: MemoryStorage }).localStorage = + new MemoryStorage(); +}); + +afterEach(() => { + (globalThis as unknown as { localStorage?: MemoryStorage }).localStorage = + undefined; +}); + +describe('persisted tool calls', () => { + it('repairs parsed argument objects while loading conversations', async () => { + localStorage.setItem( + CONVERSATIONS_KEY, + JSON.stringify({ + version: 1, + activeId: 'conversation-1', + conversations: { + 'conversation-1': { + id: 'conversation-1', + title: 'Broken chat', + createdAt: 1, + updatedAt: 1, + model: 'test-model', + messages: [ + { + id: 'assistant-1', + role: 'assistant', + content: '', + timestamp: 1, + toolCalls: [ + { + id: 'call-1', + tool: 'web_search', + arguments: { query: 'python' }, + status: 'success', + }, + ], + }, + ], + }, + }, + }), + ); + + const { useAppStore } = await import('./store'); + + expect(useAppStore.getState().messages[0].toolCalls?.[0].arguments).toBe( + '{"query":"python"}', + ); + const repaired = JSON.parse(localStorage.getItem(CONVERSATIONS_KEY) ?? '{}'); + expect( + repaired.conversations['conversation-1'].messages[0].toolCalls[0].arguments, + ).toBe('{"query":"python"}'); + }); + + it('keeps repaired conversations in memory when writeback fails', async () => { + localStorage.setItem( + CONVERSATIONS_KEY, + JSON.stringify({ + version: 1, + activeId: 'conversation-1', + conversations: { + 'conversation-1': { + id: 'conversation-1', + title: 'Readable chat', + createdAt: 1, + updatedAt: 1, + model: 'test-model', + messages: [ + { + id: 'assistant-1', + role: 'assistant', + content: '', + timestamp: 1, + toolCalls: [ + { + id: 'call-1', + tool: 'web_search', + arguments: { query: 'python' }, + status: 'success', + }, + ], + }, + ], + }, + }, + }), + ); + vi.spyOn(localStorage, 'setItem').mockImplementation(() => { + throw new DOMException('Storage quota exceeded', 'QuotaExceededError'); + }); + + const { useAppStore } = await import('./store'); + + expect(useAppStore.getState().messages).toHaveLength(1); + expect(useAppStore.getState().messages[0].toolCalls?.[0].arguments).toBe( + '{"query":"python"}', + ); + }); +}); diff --git a/frontend/src/lib/store.ts b/frontend/src/lib/store.ts index dc5909c4..f2c91f79 100644 --- a/frontend/src/lib/store.ts +++ b/frontend/src/lib/store.ts @@ -16,6 +16,7 @@ import type { } from '../types'; import type { ManagedAgent } from './api'; import { isEmbedOnlyModel } from './model-capabilities'; +import { serializeToolCallArguments } from './tool-call'; export interface CachedConnector { connector_id: string; @@ -55,7 +56,30 @@ function loadConversations(): ConversationStore { const raw = localStorage.getItem(CONVERSATIONS_KEY); if (!raw) return { version: 1, conversations: {}, activeId: null }; const parsed = JSON.parse(raw); - if (parsed.version === 1) return parsed; + if (parsed.version === 1) { + let repaired = false; + for (const conversation of Object.values(parsed.conversations ?? {}) as Conversation[]) { + for (const message of conversation.messages ?? []) { + for (const toolCall of message.toolCalls ?? []) { + const argumentsText = serializeToolCallArguments(toolCall.arguments); + if (argumentsText !== toolCall.arguments) { + toolCall.arguments = argumentsText; + repaired = true; + } + } + } + } + if (repaired) { + try { + localStorage.setItem(CONVERSATIONS_KEY, JSON.stringify(parsed)); + } catch { + // Keep the repaired conversations usable in memory when storage is + // read-only or full. A failed best-effort writeback must not make + // otherwise readable conversation history disappear from the UI. + } + } + return parsed; + } return { version: 1, conversations: {}, activeId: null }; } catch { return { version: 1, conversations: {}, activeId: null }; diff --git a/frontend/src/lib/tool-call.test.ts b/frontend/src/lib/tool-call.test.ts new file mode 100644 index 00000000..3540780a --- /dev/null +++ b/frontend/src/lib/tool-call.test.ts @@ -0,0 +1,22 @@ +import { describe, expect, it } from 'vitest'; + +import { serializeToolCallArguments } from './tool-call'; + +describe('serializeToolCallArguments', () => { + it('preserves JSON strings', () => { + expect(serializeToolCallArguments('{"query":"python"}')).toBe( + '{"query":"python"}', + ); + }); + + it('serializes parsed argument objects', () => { + expect(serializeToolCallArguments({ query: 'python' })).toBe( + '{"query":"python"}', + ); + }); + + it('uses an empty string for missing arguments', () => { + expect(serializeToolCallArguments(null)).toBe(''); + expect(serializeToolCallArguments(undefined)).toBe(''); + }); +}); diff --git a/frontend/src/lib/tool-call.ts b/frontend/src/lib/tool-call.ts new file mode 100644 index 00000000..22321cd8 --- /dev/null +++ b/frontend/src/lib/tool-call.ts @@ -0,0 +1,11 @@ +/** Convert tool-call arguments from API or persisted data into display-safe text. */ +export function serializeToolCallArguments(value: unknown): string { + if (typeof value === 'string') return value; + if (value == null) return ''; + + try { + return JSON.stringify(value) ?? String(value); + } catch { + return String(value); + } +} diff --git a/src/openjarvis/agents/_stubs.py b/src/openjarvis/agents/_stubs.py index 3cbf8dbd..3c4ce9e4 100644 --- a/src/openjarvis/agents/_stubs.py +++ b/src/openjarvis/agents/_stubs.py @@ -155,6 +155,9 @@ class BaseAgent(ABC): conversation messages, and finally the user input. """ messages: list[Message] = [] + context_messages = ( + list(context.conversation.messages) if context is not None else [] + ) # Check if the context already supplies a system message _context_has_system = ( context @@ -176,9 +179,28 @@ class BaseAgent(ABC): except Exception: effective_system_prompt = None if effective_system_prompt: + context_system_text = "\n\n".join( + message.text + for message in context_messages + if message.role == Role.SYSTEM + and message.metadata.get("memory_context") + and message.text + ) + if context_system_text: + effective_system_prompt = ( + f"{effective_system_prompt}\n\n{context_system_text}" + ) + context_messages = [ + message + for message in context_messages + if not ( + message.role == Role.SYSTEM + and message.metadata.get("memory_context") + ) + ] messages.append(Message(role=Role.SYSTEM, content=effective_system_prompt)) - if context and context.conversation.messages: - messages.extend(context.conversation.messages) + if context_messages: + messages.extend(context_messages) messages.append(Message(role=Role.USER, content=input)) return messages diff --git a/src/openjarvis/cli/ask.py b/src/openjarvis/cli/ask.py index b10fc88a..3f6152cd 100644 --- a/src/openjarvis/cli/ask.py +++ b/src/openjarvis/cli/ask.py @@ -248,6 +248,17 @@ def _get_memory_backend(config): return None +def _get_memory_facts(config): + """Load facts captured by the automatic memory service.""" + try: + from openjarvis.memory import load_configured_facts + + return load_configured_facts(config) + except Exception as exc: + logger.debug("Automatic memory facts unavailable (optional): %s", exc) + return [] + + _MEMORY_TOOLS = frozenset( {"retrieval", "memory_store", "memory_search", "memory_index", "memory_retrieve"} ) @@ -416,7 +427,8 @@ def _run_agent( from openjarvis.tools.storage.context import ContextConfig, inject_context backend = _get_memory_backend(config) - if backend is not None: + facts = _get_memory_facts(config) + if backend is not None or facts: ctx_cfg = ContextConfig( top_k=config.memory.context_top_k, min_score=config.memory.context_min_score, @@ -427,6 +439,7 @@ def _run_agent( [], backend, config=ctx_cfg, + facts=facts, ) for msg in context_messages: ctx.conversation.add(msg) @@ -963,7 +976,8 @@ def ask( ) backend = _get_memory_backend(config) - if backend is not None: + facts = _get_memory_facts(config) + if backend is not None or facts: ctx_cfg = ContextConfig( top_k=config.memory.context_top_k, min_score=config.memory.context_min_score, @@ -974,6 +988,7 @@ def ask( messages, backend, config=ctx_cfg, + facts=facts, ) except Exception as exc: logger.debug("Failed to inject memory context: %s", exc) diff --git a/src/openjarvis/cli/chat_cmd.py b/src/openjarvis/cli/chat_cmd.py index 796b8c93..c13943bc 100644 --- a/src/openjarvis/cli/chat_cmd.py +++ b/src/openjarvis/cli/chat_cmd.py @@ -2,6 +2,7 @@ from __future__ import annotations +import logging import sys from typing import List, Optional @@ -15,6 +16,8 @@ from openjarvis.core.events import EventBus from openjarvis.core.types import Message, Role from openjarvis.memory import publish_completed_exchange +logger = logging.getLogger(__name__) + def _read_input(prompt: str = "You> ") -> Optional[str]: """Read user input with graceful EOF handling.""" @@ -194,6 +197,15 @@ def chat( console.print(f"[yellow]Memory service unavailable: {exc}[/yellow]") memory_service = None + # The document backend and automatic fact store are separate persistence + # mechanisms. Context injection combines both at read time so facts from + # previous sessions are immediately available without a manual index step. + memory_backend = None + if config.agent.context_from_memory: + from openjarvis.cli.ask import _get_memory_backend + + memory_backend = _get_memory_backend(config) + # Conversation state if not system_prompt: from openjarvis.prompt.builder import SystemPromptBuilder @@ -262,15 +274,55 @@ def chat( # Add user message history.append(Message(role=Role.USER, content=user_input)) - # Generate response + generation_history = history + agent_context_message = None + if config.agent.context_from_memory: + try: + from openjarvis.memory import load_configured_facts + from openjarvis.tools.storage.context import ( + ContextConfig, + inject_context, + ) + + if memory_service is not None and hasattr(memory_service, "list_facts"): + facts = memory_service.list_facts() + else: + facts = load_configured_facts(config) + ctx_cfg = ContextConfig( + top_k=config.memory.context_top_k, + min_score=config.memory.context_min_score, + max_context_tokens=config.memory.context_max_tokens, + ) + context_messages = inject_context( + user_input, + [] if agent is not None else history, + memory_backend, + config=ctx_cfg, + facts=facts, + ) + if agent is not None: + if context_messages: + agent_context_message = context_messages[0] + else: + generation_history = context_messages + except Exception: + logger.debug("Failed to inject memory context", exc_info=True) + + # Generate response even when optional memory context is unavailable. try: if agent is not None: - response = agent.run(user_input) + agent_context = None + if agent_context_message is not None: + from openjarvis.agents._stubs import AgentContext + + agent_context = AgentContext() + agent_context.conversation.add(agent_context_message) + response = agent.run(user_input, context=agent_context) content = ( response.content if hasattr(response, "content") else str(response) ) else: - result = engine.generate(history, model=model) + result = engine.generate(generation_history, model=model) content = ( result.get("content", "") if isinstance(result, dict) diff --git a/src/openjarvis/cli/serve.py b/src/openjarvis/cli/serve.py index a9e346a3..9156fa92 100644 --- a/src/openjarvis/cli/serve.py +++ b/src/openjarvis/cli/serve.py @@ -25,6 +25,30 @@ from openjarvis.intelligence import ( logger = logging.getLogger(__name__) +_DEFAULT_TOOLS = frozenset({"think", "calculator", "web_search"}) + + +def _resolve_allowed_tools(config: object) -> tuple[set[str], bool]: + """Return configured tool names and whether the selection was explicit. + + ``tools.enabled`` is the canonical setting used by ``SystemBuilder`` and + the interactive CLI. ``agent.tools`` remains as a backward-compatible + fallback, followed by the server's default tool set when neither is set. + """ + configured = config.tools.enabled or config.agent.tools + if not configured: + return set(_DEFAULT_TOOLS), False + + if isinstance(configured, list): + allowed = { + tool.strip() + for tool in configured + if isinstance(tool, str) and tool.strip() + } + else: + allowed = {tool.strip() for tool in configured.split(",") if tool.strip()} + return allowed, True + def _unique_model_ids(model_ids: list[str]) -> list[str]: """Return model ids in first-seen order without duplicates.""" @@ -96,7 +120,7 @@ def _resolve_server_model( "--agent", "agent_name", default=None, - help="Agent for non-streaming requests (simple, orchestrator, react, openhands).", + help="Agent for chat requests (simple, orchestrator, react, openhands).", ) @click.pass_context def serve( @@ -305,21 +329,7 @@ def serve( from openjarvis.core.registry import ToolRegistry from openjarvis.tools._stubs import BaseTool - _DEFAULT_TOOLS = {"think", "calculator", "web_search"} - configured = config.agent.tools - if configured: - if isinstance(configured, list): - allowed = { - t.strip() - for t in configured - if isinstance(t, str) and t.strip() - } - else: - allowed = { - t.strip() for t in configured.split(",") if t.strip() - } - else: - allowed = _DEFAULT_TOOLS + allowed, tools_configured = _resolve_allowed_tools(config) tools = [] for name in ToolRegistry.keys(): @@ -336,7 +346,7 @@ def serve( # MCP server tools from config.tools.mcp.servers # (#461 — these were silently dropped). mcp_tools = managed_mcp_tools - if configured: + if tools_configured: mcp_tools = [ tool for tool in managed_mcp_tools @@ -406,23 +416,7 @@ def serve( from openjarvis.core.registry import ToolRegistry from openjarvis.tools._stubs import BaseTool - _DEFAULT_TOOLS = {"think", "calculator", "web_search"} - configured = config.agent.tools - if configured: - if isinstance(configured, list): - _allowed = { - t.strip() - for t in configured - if isinstance(t, str) and t.strip() - } - else: - _allowed = { - t.strip() - for t in configured.split(",") - if t.strip() - } - else: - _allowed = _DEFAULT_TOOLS + _allowed, _tools_configured = _resolve_allowed_tools(config) for _tname in ToolRegistry.keys(): if _tname not in _allowed: @@ -436,7 +430,7 @@ def serve( # Reuse the process-owned MCP pool so channels do not # open a second transport to every configured server. _ch_mcp_tools = managed_mcp_tools - if configured: + if _tools_configured: _ch_mcp_tools = [ tool for tool in managed_mcp_tools diff --git a/src/openjarvis/memory/__init__.py b/src/openjarvis/memory/__init__.py index 8e978e89..a52b9881 100644 --- a/src/openjarvis/memory/__init__.py +++ b/src/openjarvis/memory/__init__.py @@ -19,6 +19,7 @@ from openjarvis.memory.store import ( FactStore, LocalFactStore, create_fact_store, + load_configured_facts, ) __all__ = [ @@ -29,5 +30,6 @@ __all__ = [ "MemoryService", "build_memory_service", "create_fact_store", + "load_configured_facts", "publish_completed_exchange", ] diff --git a/src/openjarvis/memory/store.py b/src/openjarvis/memory/store.py index a8a63ce7..2841043a 100644 --- a/src/openjarvis/memory/store.py +++ b/src/openjarvis/memory/store.py @@ -16,7 +16,7 @@ import time from abc import ABC, abstractmethod from dataclasses import asdict, dataclass from pathlib import Path -from typing import Iterable, List +from typing import Any, Iterable, List from openjarvis.core.paths import get_config_dir from openjarvis.core.registry import FactStoreRegistry @@ -205,4 +205,30 @@ def create_fact_store( return FactStoreRegistry.create(key, path, max_facts=max_facts) -__all__ = ["Fact", "FactStore", "LocalFactStore", "create_fact_store"] +def load_configured_facts(config: Any) -> List[Fact]: + """Load automatic-memory facts from *config* when the service is enabled. + + Context injection is also used by short-lived commands such as + ``jarvis ask``, where no :class:`MemoryService` instance exists. This + helper gives those callers the same configured fact-store view without + coupling them to the service lifecycle. + """ + memory = getattr(config, "memory", None) + if memory is None or not getattr(memory, "enabled", False): + return [] + + store = create_fact_store( + getattr(memory, "backend", "local"), + path=getattr(memory, "facts_path", None), + max_facts=getattr(memory, "max_facts", 1000), + ) + return store.list() + + +__all__ = [ + "Fact", + "FactStore", + "LocalFactStore", + "create_fact_store", + "load_configured_facts", +] diff --git a/src/openjarvis/sdk.py b/src/openjarvis/sdk.py index 81d58a1f..f9f7400b 100644 --- a/src/openjarvis/sdk.py +++ b/src/openjarvis/sdk.py @@ -522,14 +522,15 @@ class Jarvis: # Context injection if context and self._config.agent.context_from_memory: try: - from openjarvis.cli.ask import _get_memory_backend + from openjarvis.cli.ask import _get_memory_backend, _get_memory_facts from openjarvis.tools.storage.context import ( ContextConfig, inject_context, ) backend = _get_memory_backend(self._config) - if backend is not None: + facts = _get_memory_facts(self._config) + if backend is not None or facts: ctx_cfg = ContextConfig( top_k=self._config.memory.context_top_k, min_score=self._config.memory.context_min_score, @@ -540,6 +541,7 @@ class Jarvis: [], backend, config=ctx_cfg, + facts=facts, ) for msg in context_messages: ctx.conversation.add(msg) @@ -570,17 +572,24 @@ class Jarvis: ) -> List[Message]: """Inject memory context into messages.""" try: - from openjarvis.cli.ask import _get_memory_backend + from openjarvis.cli.ask import _get_memory_backend, _get_memory_facts from openjarvis.tools.storage.context import ContextConfig, inject_context backend = _get_memory_backend(self._config) - if backend is not None: + facts = _get_memory_facts(self._config) + if backend is not None or facts: ctx_cfg = ContextConfig( top_k=self._config.memory.context_top_k, min_score=self._config.memory.context_min_score, max_context_tokens=self._config.memory.context_max_tokens, ) - return inject_context(query, messages, backend, config=ctx_cfg) + return inject_context( + query, + messages, + backend, + config=ctx_cfg, + facts=facts, + ) except Exception as exc: logger.warning("Failed to inject memory context: %s", exc) return messages diff --git a/src/openjarvis/server/routes.py b/src/openjarvis/server/routes.py index b07ed38b..59177e62 100644 --- a/src/openjarvis/server/routes.py +++ b/src/openjarvis/server/routes.py @@ -11,7 +11,7 @@ from fastapi import APIRouter, HTTPException, Request from fastapi.responses import StreamingResponse from openjarvis.core.paths import get_config_dir -from openjarvis.core.types import Message, Role +from openjarvis.core.types import Message, Role, ToolCall from openjarvis.server.model_capabilities import is_embed_only_model from openjarvis.server.models import ( ChatCompletionChunk, @@ -40,6 +40,15 @@ def _to_messages(chat_messages) -> list[Message]: role=role, content=m.content or "", name=m.name, + tool_calls=[ + ToolCall( + id=tool_call.get("id", ""), + name=tool_call.get("function", {}).get("name", ""), + arguments=tool_call.get("function", {}).get("arguments", "{}"), + ) + for tool_call in (m.tool_calls or []) + ] + or None, tool_call_id=m.tool_call_id, ) ) @@ -114,13 +123,15 @@ async def chat_completions(request_body: ChatCompletionRequest, request: Request memory_backend = getattr(request.app.state, "memory_backend", None) if ( config is not None - and memory_backend is not None and config.agent.context_from_memory and request_body.messages ): try: from openjarvis.tools.storage.context import ContextConfig, inject_context + memory_service = getattr(request.app.state, "memory_service", None) + facts = memory_service.list_facts() if memory_service is not None else [] + # Extract query from the last user message query_text = "" for m in reversed(request_body.messages): @@ -130,6 +141,7 @@ async def chat_completions(request_body: ChatCompletionRequest, request: Request if query_text: messages = _to_messages(request_body.messages) + messages = _ensure_identity_prompt(messages, config) ctx_cfg = ContextConfig( top_k=config.memory.context_top_k, min_score=config.memory.context_min_score, @@ -140,22 +152,35 @@ async def chat_completions(request_body: ChatCompletionRequest, request: Request messages, memory_backend, config=ctx_cfg, + facts=facts, ) - # Rebuild request messages from enriched Message objects - if len(enriched) > len(messages): - from openjarvis.server.models import ChatMessage + # Rebuild after identity/context merging so downstream engine + # adapters always receive exactly one system message. + from openjarvis.server.models import ChatMessage - new_msgs = [] - for msg in enriched: - new_msgs.append( - ChatMessage( - role=msg.role.value, - content=msg.content, - name=msg.name, - tool_call_id=getattr(msg, "tool_call_id", None), - ) + new_msgs = [] + for msg in enriched: + new_msgs.append( + ChatMessage( + role=msg.role.value, + content=msg.content, + name=msg.name, + tool_calls=[ + { + "id": tool_call.id, + "type": "function", + "function": { + "name": tool_call.name, + "arguments": tool_call.arguments, + }, + } + for tool_call in (msg.tool_calls or []) + ] + or None, + tool_call_id=getattr(msg, "tool_call_id", None), ) - request_body.messages = new_msgs + ) + request_body.messages = new_msgs except Exception: logging.getLogger("openjarvis.server").debug( "Memory context injection failed", @@ -200,12 +225,14 @@ async def chat_completions(request_body: ChatCompletionRequest, request: Request # When the client passes `tools`, stream the model's raw # OpenAI-compat function-calling decision directly from the engine # (bypassing the agent) — the streaming mirror of the non-streaming - # #454 fix. Routing tools through the agent stream bridge ignored - # `request_body.tools`, ran the agent's own tool loop, and - # word-split generic filler content into fake token deltas, so the - # caller's tool_calls were dropped entirely (the streaming analog of - # #414). For plain chat (no tools), stream token-by-token directly - # from the engine for true real-time output. + # #454 fix. Routing client-supplied tools through a server-side agent + # would execute the agent's different tool set and drop the raw tool + # call the caller expects (#414). + # + # Without client-supplied tools, keep streaming requests on the + # configured server agent so its server-side tool loop is available + # to the desktop UI and other stream:true clients (#735). Fall back to + # direct token streaming when no tool-bearing agent is configured. if request_body.tools: return await _handle_stream_tools( engine, @@ -216,6 +243,16 @@ async def chat_completions(request_body: ChatCompletionRequest, request: Request bus=getattr(request.app.state, "bus", None), memory_service=getattr(request.app.state, "memory_service", None), ) + if agent is not None and getattr(agent, "_tools", None): + return await _handle_agent_stream( + agent, + model, + request_body, + complexity_info, + trace_store=getattr(request.app.state, "trace_store", None), + bus=getattr(request.app.state, "bus", None), + memory_service=getattr(request.app.state, "memory_service", None), + ) return await _handle_stream( engine, model, @@ -547,6 +584,114 @@ def _handle_agent( ) +async def _handle_agent_stream( + agent, + model: str, + req: ChatCompletionRequest, + complexity_info=None, + *, + trace_store=None, + bus=None, + memory_service=None, +): + """Run the configured agent and return its result as an SSE response. + + Agents own the tool-execution loop, which is synchronous today. Run that + loop in a worker thread and stream its final answer once complete. This + keeps ``stream:true`` clients (including the desktop UI) on the same agent + and configured toolkit as non-streaming requests instead of bypassing the + agent and silently dropping server-side tools. + + Requests that explicitly supply OpenAI ``tools`` continue to use + ``_handle_stream_tools`` so their raw tool-call deltas are preserved. + """ + chunk_id = f"chatcmpl-{uuid.uuid4().hex[:12]}" + query_text = "" + for message in reversed(req.messages): + if message.role == "user" and message.content: + query_text = message.content + break + + async def generate(): + first_chunk = ChatCompletionChunk( + id=chunk_id, + model=model, + choices=[StreamChoice(delta=DeltaMessage(role="assistant"))], + ) + yield f"data: {first_chunk.model_dump_json()}\n\n" + + try: + response = await asyncio.to_thread( + _handle_agent, + agent, + model, + req, + complexity_info, + trace_store=trace_store, + bus=bus, + ) + except Exception as exc: + logging.getLogger("openjarvis.server").error( + "Agent stream error: %s", + exc, + exc_info=True, + ) + error_chunk = ChatCompletionChunk( + id=chunk_id, + model=model, + choices=[ + StreamChoice( + delta=DeltaMessage( + content=f"Sorry, an error occurred: {exc}", + ), + finish_reason="stop", + ) + ], + ) + yield f"data: {error_chunk.model_dump_json()}\n\n" + yield "data: [DONE]\n\n" + return + + content = _response_content(response) + if content: + content_chunk = ChatCompletionChunk( + id=chunk_id, + model=model, + choices=[StreamChoice(delta=DeltaMessage(content=content))], + ) + yield f"data: {content_chunk.model_dump_json()}\n\n" + + import json as _json + + finish_chunk = ChatCompletionChunk( + id=chunk_id, + model=model, + choices=[ + StreamChoice(delta=DeltaMessage(), finish_reason="stop"), + ], + ) + finish_data = _json.loads(finish_chunk.model_dump_json()) + finish_data["usage"] = response.usage.model_dump() + if complexity_info is not None: + finish_data["complexity"] = complexity_info.model_dump() + yield f"data: {_json.dumps(finish_data)}\n\n" + + _record_completed_exchange( + memory_service, + query_text, + content, + bus=bus, + source="server.chat.stream", + ) + yield "data: [DONE]\n\n" + + return StreamingResponse( + generate(), + media_type="text/event-stream", + headers={"Cache-Control": "no-cache", "Connection": "keep-alive"}, + ) + + async def _handle_stream_tools( engine, model: str, @@ -690,11 +835,10 @@ async def _handle_stream( ): """Stream response using SSE format. - This path streams straight from the engine, bypassing the agent / + This no-agent fallback streams straight from the engine, bypassing the ``TraceCollector``. When *trace_store* is set we accumulate the streamed tokens and record a minimal ``Trace`` once the stream completes - successfully — otherwise streamed chats (the desktop GUI's main path) - would never populate ``traces.db``. + successfully. """ import time diff --git a/src/openjarvis/server/stream_bridge.py b/src/openjarvis/server/stream_bridge.py index efe2f88a..79baf2e5 100644 --- a/src/openjarvis/server/stream_bridge.py +++ b/src/openjarvis/server/stream_bridge.py @@ -110,6 +110,12 @@ class AgentStreamBridge: def _format_named_event(self, name: str, data: dict) -> str: """Format an SSE event with an explicit ``event:`` field.""" + if name == "tool_call_start" and not isinstance(data.get("arguments"), str): + # The in-process event bus uses parsed arguments for trace/eval + # consumers, while the web SSE contract expects their JSON text. + # Copy before normalizing so other subscribers keep the object. + data = dict(data) + data["arguments"] = json.dumps(data.get("arguments")) return f"event: {name}\ndata: {json.dumps(data)}\n\n" def _run_agent(self) -> object: diff --git a/src/openjarvis/system/orchestrator.py b/src/openjarvis/system/orchestrator.py index 8570871c..6b7073f5 100644 --- a/src/openjarvis/system/orchestrator.py +++ b/src/openjarvis/system/orchestrator.py @@ -40,8 +40,9 @@ class QueryOrchestrator: messages = [Message(role=Role.USER, content=query)] - if context and s.memory_backend and s.config.agent.context_from_memory: + if context and s.config.agent.context_from_memory: try: + from openjarvis.memory import load_configured_facts from openjarvis.tools.storage.context import ( ContextConfig, inject_context, @@ -52,11 +53,13 @@ class QueryOrchestrator: min_score=s.config.memory.context_min_score, max_context_tokens=s.config.memory.context_max_tokens, ) + facts = load_configured_facts(s.config) messages = inject_context( query, messages, s.memory_backend, config=ctx_cfg, + facts=facts, ) except Exception as exc: logger.warning("Failed to inject memory context: %s", exc) diff --git a/src/openjarvis/tools/storage/context.py b/src/openjarvis/tools/storage/context.py index c8091da6..16515e04 100644 --- a/src/openjarvis/tools/storage/context.py +++ b/src/openjarvis/tools/storage/context.py @@ -2,13 +2,16 @@ from __future__ import annotations -from dataclasses import dataclass -from typing import List, Optional +from dataclasses import dataclass, replace +from typing import TYPE_CHECKING, List, Optional, Sequence from openjarvis.core.events import EventType, get_event_bus from openjarvis.core.types import Message, Role from openjarvis.tools.storage._stubs import MemoryBackend, RetrievalResult +if TYPE_CHECKING: + from openjarvis.memory.store import Fact + @dataclass(slots=True) class ContextConfig: @@ -46,28 +49,75 @@ def format_context(results: List[RetrievalResult]) -> str: def build_context_message( results: List[RetrievalResult], + facts: Sequence[Fact] = (), ) -> Message: """Create a system message with formatted context.""" - context_text = format_context(results) - content = ( - "The following context was retrieved from the knowledge" - " base. Use it to inform your response, citing sources" - " where applicable:\n\n" + context_text + sections = [] + if facts: + fact_text = "\n".join(f"- {fact.text}" for fact in facts) + sections.append( + "The following durable facts were remembered from prior " + "conversations. Use them when relevant to the user's request:\n\n" + + fact_text + ) + if results: + sections.append( + "The following context was retrieved from the knowledge" + " base. Use it to inform your response, citing sources" + " where applicable:\n\n" + format_context(results) + ) + content = "\n\n".join(sections) + return Message( + role=Role.SYSTEM, + content=content, + metadata={"memory_context": True}, ) - return Message(role=Role.SYSTEM, content=content) + + +def _merge_context_message( + messages: List[Message], + context_message: Message, +) -> List[Message]: + """Return a copy with context folded into the existing system prompt.""" + system_messages = [message for message in messages if message.role == Role.SYSTEM] + if not system_messages: + return [context_message, *messages] + + content = "\n\n".join( + part + for part in ( + *(message.text for message in system_messages), + context_message.text, + ) + if part + ) + combined = replace(system_messages[0], content=content) + merged: List[Message] = [] + inserted = False + for message in messages: + if message.role == Role.SYSTEM: + if not inserted: + merged.append(combined) + inserted = True + continue + merged.append(message) + return merged def inject_context( query: str, messages: List[Message], - backend: MemoryBackend, + backend: Optional[MemoryBackend], *, config: Optional[ContextConfig] = None, + facts: Sequence[Fact] = (), ) -> List[Message]: """Retrieve relevant context and prepend it to *messages*. Returns a **new** list — the original list is not mutated. - If no results pass the score threshold, returns the original + Automatic-memory facts are included independently of the retrieval + backend, so persisted facts remain recallable even when the document + store is empty. If no facts or results are available, returns the original messages unchanged. Parameters @@ -77,33 +127,55 @@ def inject_context( messages: The existing message list. backend: - The memory backend to search. + The memory backend to search, or ``None`` when only facts are available. config: Context injection settings (uses defaults if ``None``). + facts: + Durable facts captured by the automatic memory service. """ cfg = config or ContextConfig() if not cfg.enabled: return messages - results = backend.retrieve(query, top_k=cfg.top_k) + results = backend.retrieve(query, top_k=cfg.top_k) if backend is not None else [] # Filter by minimum score results = [r for r in results if r.score >= cfg.min_score] - if not results: - return messages - - # Truncate to max_context_tokens - truncated: List[RetrievalResult] = [] + # When both sources have data, cap facts at half the total budget so they + # cannot starve query-specific document retrieval. Unused fact budget is + # still available to documents. Newest facts win within the fact budget. + fact_budget = cfg.max_context_tokens + if results: + fact_budget //= 2 + selected_facts: List[Fact] = [] total_tokens = 0 + for fact in reversed(facts): + tokens = _count_tokens(fact.text) + if total_tokens + tokens > fact_budget: + continue + selected_facts.append(fact) + total_tokens += tokens + + # Fill the remaining context budget with retrieved documents. + truncated: List[RetrievalResult] = [] for r in results: tokens = _count_tokens(r.content) + if total_tokens + tokens > cfg.max_context_tokens: + # A large top result should not disappear solely because facts + # consumed their reserved share. Prefer that result when it fits + # the total budget on its own. + if not truncated and selected_facts and tokens <= cfg.max_context_tokens: + selected_facts = [] + total_tokens = 0 + else: + break if total_tokens + tokens > cfg.max_context_tokens: break truncated.append(r) total_tokens += tokens - if not truncated: + if not selected_facts and not truncated: return messages # Publish event @@ -114,13 +186,14 @@ def inject_context( "context_injection": True, "query": query, "num_results": len(truncated), + "num_facts": len(selected_facts), "total_tokens": total_tokens, }, ) # Build context message and prepend - ctx_msg = build_context_message(truncated) - return [ctx_msg] + list(messages) + ctx_msg = build_context_message(truncated, selected_facts) + return _merge_context_message(messages, ctx_msg) __all__ = [ diff --git a/tests/agents/test_base_agent.py b/tests/agents/test_base_agent.py index 4235959c..7c04289c 100644 --- a/tests/agents/test_base_agent.py +++ b/tests/agents/test_base_agent.py @@ -205,6 +205,40 @@ class TestBuildMessages: assert messages[1].content == "prev" assert messages[2].content == "new" + def test_prompt_builder_merges_context_system_message(self): + engine = MagicMock() + prompt_builder = MagicMock() + prompt_builder.build.return_value = "You are OpenJarvis." + agent = _ConcreteAgent(engine, "m", prompt_builder=prompt_builder) + conv = Conversation() + conv.add( + Message( + role=Role.SYSTEM, + content="Remember: user likes jazz.", + metadata={"memory_context": True}, + ) + ) + ctx = AgentContext(conversation=conv) + + messages = agent._build_messages("new", ctx) + + system_messages = [m for m in messages if m.role == Role.SYSTEM] + assert len(system_messages) == 1 + assert "You are OpenJarvis." in system_messages[0].content + assert "user likes jazz" in system_messages[0].content + + def test_prompt_builder_preserves_caller_system_context(self): + engine = MagicMock() + prompt_builder = MagicMock() + prompt_builder.build.return_value = "Agent instructions." + agent = _ConcreteAgent(engine, "m", prompt_builder=prompt_builder) + conv = Conversation() + conv.add(Message(role=Role.SYSTEM, content="You are helpful.")) + + messages = agent._build_messages("new", AgentContext(conversation=conv)) + + assert any(message.content == "You are helpful." for message in messages) + class TestGenerate: def test_delegates_to_engine(self): diff --git a/tests/cli/test_chat_cmd.py b/tests/cli/test_chat_cmd.py index 648fb833..2989e91d 100644 --- a/tests/cli/test_chat_cmd.py +++ b/tests/cli/test_chat_cmd.py @@ -18,6 +18,7 @@ from openjarvis.core.config import JarvisConfig from openjarvis.core.events import Event, EventBus, EventType from openjarvis.core.registry import AgentRegistry, ToolRegistry from openjarvis.core.types import ToolCall, ToolResult +from openjarvis.memory.store import LocalFactStore from openjarvis.tools._stubs import BaseTool, ToolSpec @@ -97,6 +98,79 @@ class TestReadInput: class TestChatAgents: + def test_direct_chat_injects_auto_memory_facts(self, tmp_path) -> None: + facts_path = tmp_path / "facts.jsonl" + LocalFactStore(facts_path).add( + "The user's favorite color is blue", + source="auto", + ) + + engine = MagicMock() + engine.engine_id = "mock" + engine.generate.return_value = {"content": "Blue."} + config = JarvisConfig() + config.intelligence.default_model = "test-model" + config.memory.enabled = True + config.memory.facts_path = str(facts_path) + config.agent.context_from_memory = True + + with ( + patch("openjarvis.cli.chat_cmd.load_config", return_value=config), + patch("openjarvis.engine.get_engine", return_value=("mock", engine)), + patch("openjarvis.intelligence.register_builtin_models"), + patch("openjarvis.memory.build_memory_service", return_value=None), + patch("openjarvis.cli.ask._get_memory_backend", return_value=None), + ): + result = CliRunner().invoke( + chat, + ["--model", "test-model"], + input="What is my favorite color?\n/quit\n", + ) + + assert result.exit_code == 0 + messages = engine.generate.call_args.args[0] + assert messages[0].role.value == "system" + assert "favorite color is blue" in messages[0].content + + def test_chat_generation_survives_fact_store_failure(self) -> None: + class _FailingMemoryService: + def start(self) -> None: + pass + + def stop(self, timeout: float = 2.0) -> None: + pass + + def list_facts(self): + raise OSError("fact store unavailable") + + engine = MagicMock() + engine.engine_id = "mock" + engine.generate.return_value = {"content": "Still working."} + config = JarvisConfig() + config.intelligence.default_model = "test-model" + config.memory.enabled = True + config.agent.context_from_memory = True + + with ( + patch("openjarvis.cli.chat_cmd.load_config", return_value=config), + patch("openjarvis.engine.get_engine", return_value=("mock", engine)), + patch("openjarvis.intelligence.register_builtin_models"), + patch( + "openjarvis.memory.build_memory_service", + return_value=_FailingMemoryService(), + ), + patch("openjarvis.cli.ask._get_memory_backend", return_value=None), + ): + result = CliRunner().invoke( + chat, + ["--model", "test-model"], + input="hello\n/quit\n", + ) + + assert result.exit_code == 0 + assert "Still working." in result.output + engine.generate.assert_called_once() + def test_simple_agent_does_not_receive_tool_only_kwargs(self) -> None: engine = MagicMock() engine.engine_id = "mock" diff --git a/tests/cli/test_serve_tools.py b/tests/cli/test_serve_tools.py new file mode 100644 index 00000000..1aedef8d --- /dev/null +++ b/tests/cli/test_serve_tools.py @@ -0,0 +1,53 @@ +"""Regression tests for tool selection during ``jarvis serve`` startup.""" + +from __future__ import annotations + +import pytest + +from openjarvis.cli.serve import _resolve_allowed_tools +from openjarvis.core.config import JarvisConfig + + +@pytest.mark.parametrize( + "configured", + [ + "code_interpreter,file_read", + ["code_interpreter", "file_read"], + ], +) +def test_tools_enabled_is_used_by_serve(configured): + config = JarvisConfig() + config.tools.enabled = configured + + allowed, explicit = _resolve_allowed_tools(config) + + assert allowed == {"code_interpreter", "file_read"} + assert explicit is True + + +def test_tools_enabled_takes_precedence_over_legacy_agent_tools(): + config = JarvisConfig() + config.tools.enabled = "file_read" + config.agent.tools = "calculator" + + allowed, explicit = _resolve_allowed_tools(config) + + assert allowed == {"file_read"} + assert explicit is True + + +def test_agent_tools_remains_a_backward_compatible_fallback(): + config = JarvisConfig() + config.agent.tools = "file_read" + + allowed, explicit = _resolve_allowed_tools(config) + + assert allowed == {"file_read"} + assert explicit is True + + +def test_serve_defaults_tools_when_no_selection_is_configured(): + allowed, explicit = _resolve_allowed_tools(JarvisConfig()) + + assert allowed == {"think", "calculator", "web_search"} + assert explicit is False diff --git a/tests/memory/test_context.py b/tests/memory/test_context.py index 8aa036ad..fe018c99 100644 --- a/tests/memory/test_context.py +++ b/tests/memory/test_context.py @@ -7,6 +7,7 @@ from typing import Any, Dict, List, Optional from openjarvis.core.events import EventBus, EventType from openjarvis.core.types import Message, Role +from openjarvis.memory.store import Fact from openjarvis.tools.storage._stubs import MemoryBackend, RetrievalResult from openjarvis.tools.storage.context import ( ContextConfig, @@ -167,6 +168,113 @@ def test_inject_context_no_results_returns_original(): assert augmented is messages +def test_inject_context_adds_auto_memory_facts_without_backend(): + messages = [Message(role=Role.USER, content="What is my favorite color?")] + facts = [Fact(text="The user's favorite color is blue", source="auto")] + + augmented = inject_context("favorite color", messages, None, facts=facts) + + assert len(augmented) == 2 + assert augmented[0].role == Role.SYSTEM + assert "remembered from prior conversations" in augmented[0].content + assert "favorite color is blue" in augmented[0].content + + +def test_inject_context_prioritizes_newest_facts_within_token_budget(): + messages = [Message(role=Role.USER, content="What do you remember?")] + facts = [ + Fact(text="old fact uses four tokens"), + Fact(text="new fact uses four tokens"), + ] + + augmented = inject_context( + "remember", + messages, + None, + config=ContextConfig(max_context_tokens=5), + facts=facts, + ) + + assert "new fact uses four tokens" in augmented[0].content + assert "old fact uses four tokens" not in augmented[0].content + + +def test_inject_context_merges_with_existing_system_message(): + messages = [ + Message(role=Role.SYSTEM, content="You are OpenJarvis."), + Message(role=Role.USER, content="What is my favorite color?"), + ] + facts = [Fact(text="The user's favorite color is blue")] + + augmented = inject_context("favorite color", messages, None, facts=facts) + + system_messages = [m for m in augmented if m.role == Role.SYSTEM] + assert len(system_messages) == 1 + assert "You are OpenJarvis." in system_messages[0].content + assert "favorite color is blue" in system_messages[0].content + assert messages[0].content == "You are OpenJarvis." + + +def test_inject_context_collapses_multiple_system_messages(): + messages = [ + Message(role=Role.SYSTEM, content="Identity."), + Message(role=Role.SYSTEM, content="Persona."), + Message(role=Role.USER, content="What do you remember?"), + ] + + augmented = inject_context( + "remember", + messages, + None, + facts=[Fact(text="User likes jazz")], + ) + + system_messages = [m for m in augmented if m.role == Role.SYSTEM] + assert len(system_messages) == 1 + assert "Identity." in system_messages[0].content + assert "Persona." in system_messages[0].content + assert "User likes jazz" in system_messages[0].content + + +def test_inject_context_reserves_budget_for_retrieved_documents(): + backend = _FakeMemory( + [RetrievalResult(content="d1 d2 d3 d4 d5", score=1.0, source="doc")] + ) + facts = [ + Fact(text="old1 old2 old3 old4 old5"), + Fact(text="new1 new2 new3 new4 new5"), + ] + + augmented = inject_context( + "query", + [Message(role=Role.USER, content="query")], + backend, + config=ContextConfig(max_context_tokens=10), + facts=facts, + ) + + assert "new1 new2 new3 new4 new5" in augmented[0].content + assert "d1 d2 d3 d4 d5" in augmented[0].content + assert "old1 old2 old3 old4 old5" not in augmented[0].content + + +def test_inject_context_prefers_large_document_that_fits_total_budget(): + backend = _FakeMemory( + [RetrievalResult(content="d1 d2 d3 d4 d5 d6 d7 d8", score=1.0)] + ) + + augmented = inject_context( + "query", + [Message(role=Role.USER, content="query")], + backend, + config=ContextConfig(max_context_tokens=10), + facts=[Fact(text="f1 f2 f3 f4 f5")], + ) + + assert "d1 d2 d3 d4 d5 d6 d7 d8" in augmented[0].content + assert "f1 f2 f3 f4 f5" not in augmented[0].content + + def test_inject_context_publishes_event(): bus = EventBus(record_history=True) results = [ diff --git a/tests/memory/test_fact_store.py b/tests/memory/test_fact_store.py index 66066c1c..149af13e 100644 --- a/tests/memory/test_fact_store.py +++ b/tests/memory/test_fact_store.py @@ -7,7 +7,11 @@ import json import pytest from openjarvis.core.registry import FactStoreRegistry -from openjarvis.memory.store import LocalFactStore, create_fact_store +from openjarvis.memory.store import ( + LocalFactStore, + create_fact_store, + load_configured_facts, +) def test_add_and_list(tmp_path): @@ -145,3 +149,28 @@ def test_create_fact_store_default_path_uses_openjarvis_home(tmp_path, monkeypat def test_create_fact_store_unknown_backend(tmp_path): with pytest.raises(ValueError): create_fact_store("cloud", path=tmp_path / "f.jsonl") + + +def test_load_configured_facts_reads_enabled_store(tmp_path): + from types import SimpleNamespace + + path = tmp_path / "facts.jsonl" + LocalFactStore(path).add("User likes jazz", source="auto") + config = SimpleNamespace( + memory=SimpleNamespace( + enabled=True, + backend="local", + facts_path=str(path), + max_facts=1000, + ) + ) + + assert [fact.text for fact in load_configured_facts(config)] == ["User likes jazz"] + + +def test_load_configured_facts_skips_disabled_memory(): + from types import SimpleNamespace + + config = SimpleNamespace(memory=SimpleNamespace(enabled=False)) + + assert load_configured_facts(config) == [] diff --git a/tests/server/test_model_management.py b/tests/server/test_model_management.py index 53017791..340ea5d3 100644 --- a/tests/server/test_model_management.py +++ b/tests/server/test_model_management.py @@ -213,6 +213,7 @@ class TestStreamingResilience: engine = _make_engine() agent = MagicMock() agent.agent_id = "simple" + agent._tools = [] agent.run.return_value = AgentResult( content="agent response", turns=1, diff --git a/tests/server/test_routes.py b/tests/server/test_routes.py index f88b0aa7..7735e224 100644 --- a/tests/server/test_routes.py +++ b/tests/server/test_routes.py @@ -11,6 +11,7 @@ fastapi = pytest.importorskip("fastapi") from fastapi.testclient import TestClient # noqa: E402 from openjarvis.core.events import EventBus, EventType # noqa: E402 +from openjarvis.core.types import Role # noqa: E402 from openjarvis.server.app import create_app # noqa: E402 # --------------------------------------------------------------------------- @@ -534,6 +535,94 @@ class TestChatCompletions: content += delta_content assert content == "Hello world" + def test_streaming_without_client_tools_uses_configured_agent(self): + """Server-side tools remain available to streaming web clients (#735).""" + from openjarvis.agents.orchestrator import OrchestratorAgent + from openjarvis.core.types import ToolResult + from openjarvis.tools._stubs import BaseTool, ToolSpec + + executions: list[str] = [] + + class _FileReadTool(BaseTool): + @property + def spec(self): + return ToolSpec( + name="file_read", + description="Read a file", + parameters={ + "type": "object", + "properties": {"path": {"type": "string"}}, + }, + ) + + def execute(self, **params): + executions.append(params["path"]) + return ToolResult( + tool_name="file_read", + content="README fixture contents", + success=True, + ) + + engine = _make_engine(content="ENGINE BYPASS") + engine.generate.side_effect = [ + { + "content": "", + "tool_calls": [ + { + "id": "call_1", + "name": "file_read", + "arguments": '{"path": "README.md"}', + } + ], + "usage": {}, + }, + { + "content": "README fixture contents", + "finish_reason": "stop", + "usage": {}, + }, + ] + agent = OrchestratorAgent( + engine, + "test-model", + tools=[_FileReadTool()], + bus=EventBus(), + max_turns=3, + temperature=0.7, + max_tokens=128, + system_prompt="Use the configured tools.", + ) + app = create_app( + engine, + "test-model", + agent=agent, + bus=EventBus(), + config=_test_config(), + ) + client = TestClient(app) + + resp = client.post( + "/v1/chat/completions", + json={ + "model": "test-model", + "messages": [{"role": "user", "content": "Read README.md"}], + "stream": True, + }, + ) + + assert resp.status_code == 200 + content = "" + for line in resp.text.strip().split("\n"): + if not line.startswith("data:") or "[DONE]" in line: + continue + data = json.loads(line[5:].strip()) + delta = data.get("choices", [{}])[0].get("delta", {}) + content += delta.get("content") or "" + + assert content == "README fixture contents" + assert executions == ["README.md"] + assert engine.generate.call_count == 2 + def test_streaming_with_tools_emits_tool_calls_and_bypasses_agent(self): """Regression for the streaming analog of #414. @@ -758,35 +847,22 @@ class TestIdentityPromptInjection: assert len(system_msgs) == 1 assert system_msgs[0].content == "Be terse." - def test_stream_uses_persona_in_single_active_inference(self, tmp_path): - """Regression for #734: web streaming must use the grounded prompt. - - The obsolete agent stream bridge ran a grounded agent inference and - then replayed the raw request through the engine, so the browser saw - an ungrounded second answer. The active endpoint must instead make a - single streaming call whose messages already include the configured - persona, even when an agent and event bus are registered. - """ - from openjarvis.core.config import MemoryFilesConfig + def test_stream_uses_grounded_agent_result_without_replay(self): + """Regression for #734: web streaming emits the agent's final answer.""" from openjarvis.core.events import EventBus - soul = tmp_path / "SOUL.md" - soul.write_text("Always introduce yourself as Jarvis Prime.") - captured: list = [] engine = _make_capturing_engine(captured) - agent = _make_agent(content="grounded agent result") - cfg = _identity_config() - cfg.memory_files = MemoryFilesConfig( - soul_path=str(soul), memory_path="", user_path="" - ) + agent = _make_agent(content="My name is Jarvis Prime.") + agent._tools = [object()] + agent._engine = engine client = TestClient( create_app( engine, "test-model", agent=agent, bus=EventBus(), - config=cfg, + config=_identity_config(), ) ) @@ -808,12 +884,9 @@ class TestIdentityPromptInjection: choices = payload.get("choices", []) if choices and choices[0]["delta"].get("content"): streamed_content += choices[0]["delta"]["content"] - assert streamed_content == "Hello world" - assert len(captured) == 1 - assert captured[0][0].role.value == "system" - assert "OpenJarvis" in captured[0][0].content - assert "Jarvis Prime" in captured[0][0].content - agent.run.assert_not_called() + assert streamed_content == "My name is Jarvis Prime." + assert captured == [] + agent.run.assert_called_once() def test_direct_injects_identity_when_absent(self): captured: list = [] @@ -855,6 +928,101 @@ class TestIdentityPromptInjection: assert len(system_msgs) == 1 assert system_msgs[0].content == "Be terse." + def test_direct_merges_identity_and_auto_memory_into_one_system_message(self): + from openjarvis.memory.store import Fact + + class _MemoryService: + def list_facts(self): + return [Fact(text="The user's favorite color is blue")] + + captured: list = [] + engine = _make_capturing_engine(captured) + cfg = _identity_config() + cfg.agent.context_from_memory = True + client = TestClient( + create_app( + engine, + "test-model", + config=cfg, + memory_service=_MemoryService(), + ) + ) + + resp = client.post( + "/v1/chat/completions", + json={ + "model": "test-model", + "messages": [{"role": "user", "content": "What is my favorite color?"}], + }, + ) + + assert resp.status_code == 200 + messages = engine.generate.call_args.args[0] + system_messages = [m for m in messages if m.role == Role.SYSTEM] + assert len(system_messages) == 1 + assert "OpenJarvis" in system_messages[0].content + assert "favorite color is blue" in system_messages[0].content + + def test_memory_context_preserves_assistant_tool_calls(self): + from openjarvis.memory.store import Fact + + class _MemoryService: + def list_facts(self): + return [Fact(text="User likes jazz")] + + captured: list = [] + engine = _make_capturing_engine(captured) + cfg = _identity_config() + cfg.agent.context_from_memory = True + client = TestClient( + create_app( + engine, + "test-model", + config=cfg, + memory_service=_MemoryService(), + ) + ) + + resp = client.post( + "/v1/chat/completions", + json={ + "model": "test-model", + "messages": [ + {"role": "user", "content": "Run the lookup"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": { + "name": "lookup", + "arguments": '{"query":"jazz"}', + }, + } + ], + }, + { + "role": "tool", + "content": "result", + "tool_call_id": "call_1", + }, + {"role": "user", "content": "What did it find?"}, + ], + }, + ) + + assert resp.status_code == 200 + messages = engine.generate.call_args.args[0] + assistant = next( + message for message in messages if message.role == Role.ASSISTANT + ) + assert assistant.tool_calls is not None + assert assistant.tool_calls[0].id == "call_1" + assert assistant.tool_calls[0].name == "lookup" + assert assistant.tool_calls[0].arguments == '{"query":"jazz"}' + def test_direct_injects_soul_persona_when_present(self, tmp_path): """Regression: /v1/chat/completions previously injected only the bare ``default_system_prompt`` blurb via a hand-rolled lookup, bypassing diff --git a/tests/server/test_stream_bridge.py b/tests/server/test_stream_bridge.py index 17d58c4f..57aa2d65 100644 --- a/tests/server/test_stream_bridge.py +++ b/tests/server/test_stream_bridge.py @@ -63,3 +63,30 @@ def test_stream_replays_grounded_agent_result_without_second_inference(): assert any(event.startswith("event: tool_results\n") for event in events) agent.run.assert_called_once() assert agent._model == "configured-model" + + +def test_tool_call_start_serializes_arguments_for_sse_without_mutating_event(): + bridge = object.__new__(AgentStreamBridge) + event_data = { + "tool": "web_search", + "arguments": {"query": "python"}, + "agent": "agent-1", + } + + event = bridge._format_named_event("tool_call_start", event_data) + payload = json.loads(event.split("data: ", 1)[1]) + + assert payload["arguments"] == '{"query": "python"}' + assert event_data["arguments"] == {"query": "python"} + + +def test_tool_call_start_preserves_already_serialized_arguments(): + bridge = object.__new__(AgentStreamBridge) + + event = bridge._format_named_event( + "tool_call_start", + {"tool": "web_search", "arguments": '{"query":"python"}'}, + ) + payload = json.loads(event.split("data: ", 1)[1]) + + assert payload["arguments"] == '{"query":"python"}'