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:
Eri Barrett 2025-06-24 15:25:33 -04:00 committed by GitHub
parent a67fb3da95
commit 80652d4395
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
5 changed files with 645 additions and 2 deletions

View File

@ -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)

View File

@ -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}"

View File

@ -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

View File

@ -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"
)

View File

@ -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"]
)