get deriver status for peer, optional session param (#132)
* Initial Model Changes * fix migration * update schemas * handle router changes * make name FK and corresponding crud changes * fix routers * comment metamessage references * add bulk peer session operations * update messages router * fix require_auth to make app runnable * remove peer from get messages * add new routes * implement new crud methods for session peers * alter keys router * add feature flags dict and token limit + fix SessionContext * fix: paginate get_session_peers and make tokens/summary query params in get_session_context * feat: add create_messages_for_peer, get_messages_for_peer * fix: make session_peers a Table * finalize upgrade * fix: working migration * fixes: schemas, crud, routes * add token count * fix migration errors discovered from db with data in it * fixes: unify with sdk * downgrade * feat: swap jwts to new paradigm * fix unit tests * fix tests pt 2 * fix: handle foreign key errors in create_messages * fix downgrade * downgrade queue changes * feat: add search to resources, make get_messages handle limits, add get_representation to peer * chore: beef up tests * fix: move chat and rep params to post body, add target to get_representation * fix get_user_protected_collection and embedding store * feat: add peer config to models, crud, schemas, routes * fix: update tests and fix list(tuple()) to dict() * add session peer left_at/joined_at and modify enqueue * [wip]: feat: refactor history to match new paradigm and implement get_context * fix messages enqueue and test it * chore: align deriver and new honcho paradigm * chore: update consumer * chore: get rid of is_user * feat: change queue tables to new key strat * fix: convert queue session_id to str properly * fix downgrade migration * feat: re-integrate old deriver * chore: coderabbit review, lots of small bug fixes * fix: fix batch migration of messages and token count * fix: mock ModelClient * CodeRabbit comments * CR comments 2 * fix: handle metadata and feature flags properly in get_or_creates * cr comments 3 * feature flag to configuration * feat: add real get crud * fix: remove reverse param from places it does not belong * add session.name constraint; narrow task type; disable deriver from configuration * get_or_add_peers_to_session + session peers limit * feat: get deriver status for peer, optional session param * fix: add internal_metadata, fix agent * rename to get_deriver_status, simplify * fix: move working rep into crud get/set, unstub get_working_representation * fix: don't payload metadata * peer protected collection -> global / local rep collections * Simplify control flow, use session_name vs id * coderabbit syntax errors * coderabbit changes * ruff formatting * Revert non-src changes from ruff formatting * move status endpoint into workspace, protects session_name, peer is optional * fix: optimize db query for deriver status * chore: add tests for queue status endpoint, add some extra validation in endpoint, reduce post-processing --------- Co-authored-by: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> Co-authored-by: Rajat Ahuja <rahuja445@gmail.com> Co-authored-by: Benjamin McCormick <docterformer@protonmail.com> Co-authored-by: doria <93405247+dr-frmr@users.noreply.github.com>
This commit is contained in:
parent
a67fb3da95
commit
80652d4395
|
|
@ -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)
|
||||
|
|
|
|||
204
src/crud.py
204
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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
)
|
||||
Loading…
Reference in New Issue