Skip to content
Merged
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
33 changes: 33 additions & 0 deletions meshmind/cli/repl.py
Original file line number Diff line number Diff line change
Expand Up @@ -116,6 +116,39 @@ async def run_mesh_repl(
async def _run_query(mesh: Mesh, text: str) -> None:
click.echo(f"{LIGHT_PURPLE}Querying…{RESET}")
try:
if hasattr(mesh, "query_stream"):
result_text = ""
nodes_used: list[str] = []
duration: float | None = None
printed_chunk = False
click.echo()
async for event in mesh.query_stream(text):
e_type = str(event.get("type", "") or "")
if e_type == "chunk":
piece = str(event.get("text", "") or "")
if piece:
printed_chunk = True
result_text += piece
click.echo(piece, nl=False)
elif e_type == "done":
nodes_used = [str(n) for n in event.get("nodes_used") or []]
duration = event.get("duration")
if not result_text:
result_text = str(event.get("result", "") or "")

if printed_chunk:
click.echo()
if not result_text:
result_text = "(empty response)"

if nodes_used:
click.echo(f"{LIGHT_PURPLE}Routing: {' -> '.join(nodes_used)}{RESET}")
if duration is not None:
click.echo(f"{LIGHT_PURPLE}Duration: {duration}s{RESET}")
if not printed_chunk:
click.echo(result_text)
return

result = await mesh.query(text)
except Exception as e:
click.echo(f"\rError: {e}", err=True)
Expand Down
174 changes: 139 additions & 35 deletions meshmind/sdk/mesh.py
Original file line number Diff line number Diff line change
Expand Up @@ -288,21 +288,11 @@ async def query(self, text: str, target: str | None = None) -> QueryResult:
if not self._running:
return QueryResult(text="Mesh is not running.", nodes_used=[])

# Find the node to query through
if target:
node = next((n for n in self._nodes if n.info.node_name == target), None)
if not node:
return QueryResult(
text=f"Node '{target}' not found in mesh.",
nodes_used=[],
)
elif self._coordinator:
node = self._coordinator
else:
# No coordinator - use first node
node = self._nodes[0] if self._nodes else None
if not node:
return QueryResult(text="No nodes available.", nodes_used=[])
node, err = self._resolve_query_node(target)
if err:
return QueryResult(text=err, nodes_used=[])
if node is None:
return QueryResult(text="No nodes available.", nodes_used=[])

logger.info(
"Query entrypoint node: %s (%s)",
Expand All @@ -313,27 +303,92 @@ async def query(self, text: str, target: str | None = None) -> QueryResult:
result_text, trace = await node.query(text)

query_result = QueryResult(text=result_text)
self._apply_trace_to_query_result(query_result, trace)

await self._emit(
"query_completed",
{
"query": text[:100],
"duration": query_result.duration,
"nodes_used": query_result.nodes_used,
},
)

if trace:
query_result.duration = round(trace.end_time - trace.start_time, 2) if trace.end_time else None
query_result.nodes_used = [r["source_node"] for r in trace.responses]
query_result.unavailable_nodes = trace.unavailable_nodes
query_result.trace = {
"correlation_id": trace.correlation_id,
"routing_duration_ms": (
int((trace.routing_completed_at - trace.routing_started_at) * 1000)
if trace.routing_started_at and trace.routing_completed_at
else None
),
"aggregation_duration_ms": (
int((trace.aggregation_completed_at - trace.aggregation_started_at) * 1000)
if trace.aggregation_started_at and trace.aggregation_completed_at
else None
),
"sub_queries": trace.sub_queries,
"responses": trace.responses,
"unavailable_nodes": trace.unavailable_nodes,
logger.info(
"Query finished via %s; workers=%s; duration=%ss",
node.info.node_name,
", ".join(query_result.nodes_used) if query_result.nodes_used else "none",
query_result.duration if query_result.duration is not None else "n/a",
)

return query_result

async def query_stream(self, text: str, target: str | None = None):
"""Send a query through the mesh and stream chunks when available."""
if not self._running:
msg = "Mesh is not running."
yield {"type": "chunk", "text": msg}
yield {
"type": "done",
"result": msg,
"duration": None,
"nodes_used": [],
"unavailable_nodes": [],
"trace": None,
}
return

node, err = self._resolve_query_node(target)
if err:
yield {"type": "chunk", "text": err}
yield {
"type": "done",
"result": err,
"duration": None,
"nodes_used": [],
"unavailable_nodes": [],
"trace": None,
}
return
if node is None:
msg = "No nodes available."
yield {"type": "chunk", "text": msg}
yield {
"type": "done",
"result": msg,
"duration": None,
"nodes_used": [],
"unavailable_nodes": [],
"trace": None,
}
return

logger.info(
"Query entrypoint node: %s (%s)",
node.info.node_name,
node.info.node_type,
)

result_text = ""
trace = None
if hasattr(node, "query_stream"):
async for event in node.query_stream(text):
e_type = str(event.get("type", "") or "")
if e_type == "chunk":
piece = str(event.get("text", "") or "")
if piece:
result_text += piece
yield {"type": "chunk", "text": piece}
elif e_type == "done":
trace = event.get("trace")
if not result_text:
result_text = str(event.get("result", "") or "")
else:
result_text, trace = await node.query(text)
yield {"type": "chunk", "text": result_text}

query_result = QueryResult(text=result_text)
self._apply_trace_to_query_result(query_result, trace)

await self._emit(
"query_completed",
Expand All @@ -351,7 +406,56 @@ async def query(self, text: str, target: str | None = None) -> QueryResult:
query_result.duration if query_result.duration is not None else "n/a",
)

return query_result
yield {
"type": "done",
"result": query_result.text,
"duration": query_result.duration,
"nodes_used": query_result.nodes_used,
"unavailable_nodes": query_result.unavailable_nodes,
"trace": query_result.trace,
}

def _resolve_query_node(self, target: str | None) -> tuple[MeshNode | None, str | None]:
"""Return the entrypoint node for a query and an optional user-facing error."""
if target:
node = next((n for n in self._nodes if n.info.node_name == target), None)
if not node:
return None, f"Node '{target}' not found in mesh."
return node, None

if self._coordinator:
return self._coordinator, None

if self._nodes:
return self._nodes[0], None

return None, None

@staticmethod
def _apply_trace_to_query_result(query_result: QueryResult, trace: Any) -> None:
"""Populate a QueryResult from an orchestrator trace object when available."""
if not trace:
return

query_result.duration = round(trace.end_time - trace.start_time, 2) if trace.end_time else None
query_result.nodes_used = [r["source_node"] for r in trace.responses]
query_result.unavailable_nodes = trace.unavailable_nodes
query_result.trace = {
"correlation_id": trace.correlation_id,
"routing_duration_ms": (
int((trace.routing_completed_at - trace.routing_started_at) * 1000)
if trace.routing_started_at and trace.routing_completed_at
else None
),
"aggregation_duration_ms": (
int((trace.aggregation_completed_at - trace.aggregation_started_at) * 1000)
if trace.aggregation_started_at and trace.aggregation_completed_at
else None
),
"sub_queries": trace.sub_queries,
"responses": trace.responses,
"unavailable_nodes": trace.unavailable_nodes,
}

def status(self) -> MeshStatus:
"""Get the current mesh status."""
Expand Down
40 changes: 40 additions & 0 deletions tests/test_cli_session_log.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,46 @@ async def _fake_read_line(prompt: str) -> str | None:
assert stop.is_set()


@pytest.mark.asyncio
async def test_repl_run_query_streams_chunks(capsys) -> None:
from meshmind.cli.repl import _run_query

class _DummyMesh:
async def query_stream(self, text: str):
assert text == "hello"
yield {"type": "chunk", "text": "Hello"}
yield {"type": "chunk", "text": " world"}
yield {
"type": "done",
"result": "Hello world",
"nodes_used": ["coordinator", "writer"],
"duration": 1.23,
}

await _run_query(_DummyMesh(), "hello")
out = capsys.readouterr().out
assert "Hello world" in out
assert "Routing: coordinator -> writer" in out
assert "Duration: 1.23s" in out


@pytest.mark.asyncio
async def test_repl_run_query_falls_back_to_non_stream(capsys) -> None:
from meshmind.cli.repl import _run_query
from meshmind.sdk.results import QueryResult

class _DummyMesh:
async def query(self, text: str) -> QueryResult:
assert text == "hello"
return QueryResult(text="fallback", nodes_used=["coordinator"], duration=0.5)

await _run_query(_DummyMesh(), "hello")
out = capsys.readouterr().out
assert "fallback" in out
assert "Routing: coordinator" in out
assert "Duration: 0.5s" in out


def test_session_log_file_path(tmp_path: Path) -> None:
cfg = tmp_path / "proj" / "meshmind.yaml"
cfg.parent.mkdir(parents=True)
Expand Down
Loading