diff --git a/migrations/versions/d429de0e5338_adopt_peer_paradigm.py b/migrations/versions/d429de0e5338_adopt_peer_paradigm.py index 04cca99c..e0d08f6f 100644 --- a/migrations/versions/d429de0e5338_adopt_peer_paradigm.py +++ b/migrations/versions/d429de0e5338_adopt_peer_paradigm.py @@ -589,6 +589,7 @@ def update_messages_table(schema: str, inspector) -> None: op.create_index( "ix_messages_workspace_name", "messages", ["workspace_name"], schema=schema ) + # Create full text search index on content column op.create_index( "idx_messages_content_gin", @@ -1531,6 +1532,7 @@ def restore_messages_table(schema: str, inspector) -> None: op.drop_index( "ix_messages_workspace_name", table_name="messages", schema=schema ) + # Drop full text search index if index_exists("messages", "idx_messages_content_gin", inspector): op.drop_index("idx_messages_content_gin", table_name="messages", schema=schema) diff --git a/src/crud.py b/src/crud.py index 74be881f..7292bbbb 100644 --- a/src/crud.py +++ b/src/crud.py @@ -890,6 +890,7 @@ async def _get_or_add_peers_to_session( ) await db.execute(stmt) + # Return all active session peers after the upsert select_stmt = select(models.SessionPeer).where( models.SessionPeer.session_name == session_name, models.SessionPeer.workspace_name == workspace_name, @@ -908,12 +909,14 @@ async def get_peer_config( """ Get the configuration for a peer in a session. + Args: db: Database session workspace_name: Name of the workspace session_name: Name of the session peer_id: Name of the peer + Returns: Configuration for the peer @@ -979,7 +982,6 @@ async def set_peer_config( session_peer.configuration["observe_me"] = config.observe_me await db.commit() - return async def search( @@ -1655,5 +1657,205 @@ async def get_duplicate_documents( return list(result.scalars().all()) # Convert to list to match the return type +######################################################## +# deriver queue methods +######################################################## + + +async def get_deriver_status( + db: AsyncSession, + workspace_name: str, + peer_name: Optional[str] = None, + session_name: Optional[str] = None, + include_sender: bool = False, +) -> schemas.DeriverStatus: + """ + Get the deriver processing status, optionally filtered by peer and/or session. + + Args: + db: Database session + workspace_name: Name of the workspace + peer_name: Optional name of the peer to filter by + session_name: Optional session name to filter by + include_sender: Whether to include work units where peer is the sender + + Returns: + DeriverStatus: Schema containing processing status + + Raises: + ValueError: If neither peer_name nor session_name is provided + """ + if (peer_name is None or peer_name == "") and ( + session_name is None or session_name == "" + ): + raise ValueError("At least one of peer_name or session_name must be provided") + + # Normalize empty strings to None for consistent handling + normalized_peer_name = peer_name if peer_name else None + normalized_session_name = session_name if session_name else None + + stmt = _build_queue_status_query( + workspace_name, normalized_peer_name, normalized_session_name, include_sender + ) + result = await db.execute(stmt) + rows = result.fetchall() + + counts = _process_queue_rows(rows) + return _build_status_response(peer_name, session_name, counts) + + +def _build_queue_status_query( + workspace_name: str, + peer_name: Optional[str], + session_name: Optional[str], + include_sender: bool, +): + """Build SQL query for queue status with validation and aggregation.""" + from sqlalchemy import case, func + + sender_name_expr = models.QueueItem.payload["sender_name"].astext + target_name_expr = models.QueueItem.payload["target_name"].astext + task_type_expr = models.QueueItem.payload["task_type"].astext + + # Define conditions for cleaner window functions + is_completed = models.QueueItem.processed + is_in_progress = (~models.QueueItem.processed) & ( + models.ActiveQueueSession.id.isnot(None) + ) + is_pending = (~models.QueueItem.processed) & ( + models.ActiveQueueSession.id.is_(None) + ) + + # Use window functions to calculate totals and per-session counts in SQL + stmt = select( + models.QueueItem.session_id, + # Overall totals using window functions + func.count().over().label("total"), + func.count(case((is_completed, 1))).over().label("completed"), + func.count(case((is_in_progress, 1))).over().label("in_progress"), + func.count(case((is_pending, 1))).over().label("pending"), + # Per-session totals using partitioned window functions + func.count() + .over(partition_by=models.QueueItem.session_id) + .label("session_total"), + func.count(case((is_completed, 1))) + .over(partition_by=models.QueueItem.session_id) + .label("session_completed"), + func.count(case((is_in_progress, 1))) + .over(partition_by=models.QueueItem.session_id) + .label("session_in_progress"), + func.count(case((is_pending, 1))) + .over(partition_by=models.QueueItem.session_id) + .label("session_pending"), + ).select_from(models.QueueItem) + + stmt = stmt.outerjoin( + models.ActiveQueueSession, + (models.QueueItem.session_id == models.ActiveQueueSession.session_id) + & (sender_name_expr == models.ActiveQueueSession.sender_name) + & (target_name_expr == models.ActiveQueueSession.target_name) + & (task_type_expr == models.ActiveQueueSession.task_type), + ) + + if peer_name is not None: + stmt = stmt.outerjoin( + models.Peer, + (models.Peer.name == peer_name) + & (models.Peer.workspace_name == workspace_name), + ) + + if session_name is not None: + stmt = stmt.outerjoin( + models.Session, + (models.Session.name == session_name) + & (models.Session.workspace_name == workspace_name), + ) + stmt = stmt.where(models.QueueItem.session_id == models.Session.id) + + if peer_name is not None: + if include_sender: + from sqlalchemy import or_ + + stmt = stmt.where( + or_( + target_name_expr == peer_name, + sender_name_expr == peer_name, + ) + ) + else: + stmt = stmt.where(target_name_expr == peer_name) + + return stmt + + +def _process_queue_rows(rows): + """Process query results that already contain aggregated counts.""" + if not rows: + return { + "total": 0, + "completed": 0, + "in_progress": 0, + "pending": 0, + "sessions": {}, + } + + # Since we're using window functions, all rows have the same overall totals + # We just need the first row for overall counts + first_row = rows[0] + + # Build sessions dictionary from unique session_ids + sessions = {} + seen_sessions = set() + + for row in rows: + if row.session_id and row.session_id not in seen_sessions: + sessions[row.session_id] = { + "completed": row.session_completed, + "in_progress": row.session_in_progress, + "pending": row.session_pending, + } + seen_sessions.add(row.session_id) + + return { + "total": first_row.total, + "completed": first_row.completed, + "in_progress": first_row.in_progress, + "pending": first_row.pending, + "sessions": sessions, + } + + +def _build_status_response( + peer_name: Optional[str], session_name: Optional[str], counts: dict +): + """Build the final response object.""" + base_response = { + "peer_id": peer_name, + "total_work_units": counts["total"], + "completed_work_units": counts["completed"], + "in_progress_work_units": counts["in_progress"], + "pending_work_units": counts["pending"], + } + + if session_name: + return schemas.DeriverStatus(session_id=session_name, **base_response) + + sessions = {} + for session_id, data in counts["sessions"].items(): + total = data["completed"] + data["in_progress"] + data["pending"] + sessions[session_id] = schemas.DeriverStatus( + peer_id=peer_name, + session_id=session_id, + total_work_units=total, + completed_work_units=data["completed"], + in_progress_work_units=data["in_progress"], + pending_work_units=data["pending"], + ) + + return schemas.DeriverStatus( + sessions=sessions if sessions else None, **base_response + ) + + def construct_collection_name(peer_name: str, target_name: str) -> str: return f"{peer_name}_{target_name}" diff --git a/src/routers/workspaces.py b/src/routers/workspaces.py index a1aa4524..1d0bc74b 100644 --- a/src/routers/workspaces.py +++ b/src/routers/workspaces.py @@ -1,7 +1,7 @@ import logging from typing import Optional -from fastapi import APIRouter, Body, Depends, Path +from fastapi import APIRouter, Body, Depends, HTTPException, Path, Query from fastapi_pagination import Page from fastapi_pagination.ext.sqlalchemy import paginate @@ -105,3 +105,40 @@ async def search_workspace( stmt = await crud.search(query, workspace_name=workspace_id) return await paginate(db, stmt) + + +@router.get( + "/{workspace_id}/deriver/status", + response_model=schemas.DeriverStatus, + dependencies=[Depends(require_auth(workspace_name="workspace_id"))], +) +async def get_deriver_status( + workspace_id: str = Path(..., description="ID of the workspace"), + peer_id: Optional[str] = Query(None, description="Optional peer ID to filter by"), + session_id: Optional[str] = Query( + None, description="Optional session ID to filter by" + ), + include_sender: bool = Query( + False, description="Include work units triggered by this peer" + ), + db=db, +): + """Get the deriver processing status, optionally scoped to a peer and/or session""" + # Validate that at least one of peer_id or session_id is provided + if peer_id is None and session_id is None: + raise HTTPException( + status_code=400, + detail="At least one of 'peer_id' or 'session_id' must be provided", + ) + + try: + return await crud.get_deriver_status( + db, + workspace_name=workspace_id, + peer_name=peer_id, + session_name=session_id, + include_sender=include_sender, + ) + except ValueError as e: + logger.warning(f"Invalid request parameters: {str(e)}") + raise HTTPException(status_code=400, detail=str(e)) from e diff --git a/src/schemas.py b/src/schemas.py index 0432de25..0978dcd6 100644 --- a/src/schemas.py +++ b/src/schemas.py @@ -225,3 +225,22 @@ class DialecticOptions(BaseModel): class DialecticResponse(BaseModel): content: str + + +class DeriverStatus(BaseModel): + peer_id: Optional[str] = Field( + default=None, + description="ID of the peer (optional when filtering by session only)", + ) + session_id: Optional[str] = Field( + default=None, description="Session ID if filtered by session" + ) + total_work_units: int = Field(description="Total work units") + completed_work_units: int = Field(description="Completed work units") + in_progress_work_units: int = Field( + description="Work units currently being processed" + ) + pending_work_units: int = Field(description="Work units waiting to be processed") + sessions: Optional[dict[str, "DeriverStatus"]] = Field( + default=None, description="Per-session status when not filtered by session" + ) diff --git a/tests/routes/test_queue_status.py b/tests/routes/test_queue_status.py new file mode 100644 index 00000000..3eae1bd1 --- /dev/null +++ b/tests/routes/test_queue_status.py @@ -0,0 +1,383 @@ +import pytest +from nanoid import generate as generate_nanoid + +from src import models + + +@pytest.mark.asyncio +class TestDeriverStatusEndpoint: + """Test suite for the /deriver/status endpoint""" + + async def test_get_deriver_status_peer_only(self, client, db_session, sample_data): + """Test getting deriver status filtered by peer only""" + test_workspace, test_peer = sample_data + + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id={test_peer.name}" + ) + assert response.status_code == 200 + data = response.json() + + # Check response structure matches DeriverStatus schema + assert "peer_id" in data + assert "total_work_units" in data + assert "completed_work_units" in data + assert "in_progress_work_units" in data + assert "pending_work_units" in data + assert data["peer_id"] == test_peer.name + assert isinstance(data["total_work_units"], int) + assert isinstance(data["completed_work_units"], int) + assert isinstance(data["in_progress_work_units"], int) + assert isinstance(data["pending_work_units"], int) + + async def test_get_deriver_status_session_only( + self, client, db_session, sample_data + ): + """Test getting deriver status filtered by session only""" + test_workspace, test_peer = sample_data + + # Create a test session + test_session = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add(test_session) + await db_session.commit() + + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?session_id={test_session.name}" + ) + assert response.status_code == 200 + data = response.json() + + # Check response structure + assert "session_id" in data + assert "total_work_units" in data + assert "completed_work_units" in data + assert "in_progress_work_units" in data + assert "pending_work_units" in data + assert data["session_id"] == test_session.name + assert isinstance(data["total_work_units"], int) + + async def test_get_deriver_status_peer_and_session( + self, client, db_session, sample_data + ): + """Test getting deriver status filtered by both peer and session""" + test_workspace, test_peer = sample_data + + # Create a test session + test_session = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add(test_session) + await db_session.commit() + + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id={test_peer.name}&session_id={test_session.name}" + ) + assert response.status_code == 200 + data = response.json() + + # Should have both peer_id and session_id in response + assert data["peer_id"] == test_peer.name + assert data["session_id"] == test_session.name + assert "total_work_units" in data + assert "completed_work_units" in data + assert "in_progress_work_units" in data + assert "pending_work_units" in data + + async def test_get_deriver_status_with_include_sender_true( + self, client, db_session, sample_data + ): + """Test getting deriver status with include_sender=True""" + test_workspace, test_peer = sample_data + + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id={test_peer.name}&include_sender=true" + ) + assert response.status_code == 200 + data = response.json() + + assert data["peer_id"] == test_peer.name + assert "total_work_units" in data + + async def test_get_deriver_status_with_include_sender_false( + self, client, db_session, sample_data + ): + """Test getting deriver status with include_sender=False (default)""" + test_workspace, test_peer = sample_data + + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id={test_peer.name}&include_sender=false" + ) + assert response.status_code == 200 + data = response.json() + + assert data["peer_id"] == test_peer.name + assert "total_work_units" in data + + async def test_get_deriver_status_no_parameters(self, client, sample_data): + """Test getting deriver status without required parameters returns 400""" + test_workspace, _ = sample_data + + response = client.get(f"/v2/workspaces/{test_workspace.name}/deriver/status") + assert response.status_code == 400 + data = response.json() + assert "detail" in data + assert ( + "At least one of 'peer_id' or 'session_id' must be provided" + in data["detail"] + ) + + async def test_get_deriver_status_nonexistent_peer(self, client, sample_data): + """Test getting deriver status for nonexistent peer returns empty result""" + test_workspace, _ = sample_data + nonexistent_peer = str(generate_nanoid()) + + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id={nonexistent_peer}" + ) + assert response.status_code == 200 + data = response.json() + assert data["peer_id"] == nonexistent_peer + assert data["total_work_units"] == 0 + assert data["completed_work_units"] == 0 + assert data["in_progress_work_units"] == 0 + assert data["pending_work_units"] == 0 + + async def test_get_deriver_status_nonexistent_session(self, client, sample_data): + """Test getting deriver status for nonexistent session returns empty result""" + test_workspace, _ = sample_data + nonexistent_session = str(generate_nanoid()) + + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?session_id={nonexistent_session}" + ) + assert response.status_code == 200 + data = response.json() + assert data["session_id"] == nonexistent_session + assert data["total_work_units"] == 0 + assert data["completed_work_units"] == 0 + assert data["in_progress_work_units"] == 0 + assert data["pending_work_units"] == 0 + + async def test_get_deriver_status_nonexistent_workspace(self, client): + """Test getting deriver status for nonexistent workspace returns empty result""" + nonexistent_workspace = str(generate_nanoid()) + fake_peer = str(generate_nanoid()) + + response = client.get( + f"/v2/workspaces/{nonexistent_workspace}/deriver/status?peer_id={fake_peer}" + ) + # This should return empty result since workspace/peer doesn't exist + assert response.status_code == 200 + data = response.json() + assert data["peer_id"] == fake_peer + assert data["total_work_units"] == 0 + assert data["completed_work_units"] == 0 + assert data["in_progress_work_units"] == 0 + assert data["pending_work_units"] == 0 + + async def test_get_deriver_status_with_queue_items( + self, client, db_session, sample_data + ): + """Test getting deriver status when there are actual queue items""" + test_workspace, test_peer = sample_data + + # Create a test session + test_session = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add(test_session) + await db_session.flush() + + # Create some queue items to test with + queue_items = [ + models.QueueItem( + session_id=test_session.id, + processed=False, + payload={ + "task_type": "representation", + "sender_name": test_peer.name, + "target_name": test_peer.name, + "workspace_name": test_workspace.name, + "session_name": test_session.name, + }, + ), + models.QueueItem( + session_id=test_session.id, + processed=True, + payload={ + "task_type": "representation", + "sender_name": test_peer.name, + "target_name": test_peer.name, + "workspace_name": test_workspace.name, + "session_name": test_session.name, + }, + ), + ] + db_session.add_all(queue_items) + await db_session.commit() + + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id={test_peer.name}&session_id={test_session.name}" + ) + assert response.status_code == 200 + data = response.json() + + # Should have some work units + assert data["total_work_units"] >= 2 + assert data["completed_work_units"] >= 1 + assert data["pending_work_units"] >= 1 + + async def test_get_deriver_status_with_sessions_breakdown( + self, client, db_session, sample_data + ): + """Test getting deriver status shows sessions breakdown when appropriate""" + test_workspace, test_peer = sample_data + + # Create multiple test sessions + test_session1 = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + test_session2 = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add_all([test_session1, test_session2]) + await db_session.flush() + + # Create queue items for different sessions + queue_items = [ + models.QueueItem( + session_id=test_session1.id, + processed=False, + payload={ + "task_type": "representation", + "sender_name": test_peer.name, + "target_name": test_peer.name, + "workspace_name": test_workspace.name, + }, + ), + models.QueueItem( + session_id=test_session2.id, + processed=True, + payload={ + "task_type": "representation", + "sender_name": test_peer.name, + "target_name": test_peer.name, + "workspace_name": test_workspace.name, + }, + ), + ] + db_session.add_all(queue_items) + await db_session.commit() + + # Get status for peer only (should include sessions breakdown) + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id={test_peer.name}" + ) + assert response.status_code == 200 + data = response.json() + + # Should have sessions breakdown when querying by peer only + if "sessions" in data and data["sessions"]: + assert isinstance(data["sessions"], dict) + # Each session should have its own status + for _, session_data in data["sessions"].items(): + assert "total_work_units" in session_data + assert "completed_work_units" in session_data + assert "in_progress_work_units" in session_data + assert "pending_work_units" in session_data + + async def test_get_deriver_status_empty_parameters(self, client, sample_data): + """Test various edge cases with empty or invalid parameters""" + test_workspace, _ = sample_data + + # Test with empty peer_id + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id=" + ) + assert response.status_code == 400 + + # Test with empty session_id + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?session_id=" + ) + assert response.status_code == 400 + + async def test_get_deriver_status_boolean_parameter_variations( + self, client, sample_data + ): + """Test different boolean parameter formats for include_sender""" + test_workspace, test_peer = sample_data + + # Test with string 'true' + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id={test_peer.name}&include_sender=true" + ) + assert response.status_code == 200 + + # Test with string 'false' + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id={test_peer.name}&include_sender=false" + ) + assert response.status_code == 200 + + # Test with boolean True + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id={test_peer.name}&include_sender=True" + ) + assert response.status_code == 200 + + # Test with boolean False + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status?peer_id={test_peer.name}&include_sender=False" + ) + assert response.status_code == 200 + + async def test_get_deriver_status_response_consistency( + self, client, db_session, sample_data + ): + """Test that response structure is consistent across different parameter combinations""" + test_workspace, test_peer = sample_data + + # Create a test session + test_session = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add(test_session) + await db_session.commit() + + # Test different parameter combinations and ensure consistent response structure + test_cases = [ + f"?peer_id={test_peer.name}", + f"?session_id={test_session.name}", + f"?peer_id={test_peer.name}&session_id={test_session.name}", + f"?peer_id={test_peer.name}&include_sender=true", + f"?session_id={test_session.name}&include_sender=false", + ] + + for params in test_cases: + response = client.get( + f"/v2/workspaces/{test_workspace.name}/deriver/status{params}" + ) + assert response.status_code == 200 + data = response.json() + + # All responses should have these base fields + assert "total_work_units" in data + assert "completed_work_units" in data + assert "in_progress_work_units" in data + assert "pending_work_units" in data + + # Verify counts are non-negative integers + assert data["total_work_units"] >= 0 + assert data["completed_work_units"] >= 0 + assert data["in_progress_work_units"] >= 0 + assert data["pending_work_units"] >= 0 + + # Verify total equals sum of components + assert data["total_work_units"] == ( + data["completed_work_units"] + + data["in_progress_work_units"] + + data["pending_work_units"] + )