fix: resolve get_effective_observe me race condition, default peer config (#176)
* fix: resolve get_effective_observe me race condition, default peer config * fix: preserve custom config even after leaving * chore: test cases, enqueue types
This commit is contained in:
parent
0664410878
commit
23557ced02
|
|
@ -425,28 +425,30 @@ async def get_peers_from_session(
|
|||
async def get_session_peer_configuration(
|
||||
workspace_name: str,
|
||||
session_name: str,
|
||||
) -> Select[tuple[str, dict[str, Any], dict[str, Any]]]:
|
||||
) -> Select[tuple[str, dict[str, Any], dict[str, Any], bool]]:
|
||||
"""
|
||||
Get configuration from both SessionPeer and Peer tables for active peers in a session.
|
||||
Get configuration from both SessionPeer and Peer tables for all peers in a session.
|
||||
NOTE: does not filter for active peers. Will return peers that have left the session.
|
||||
|
||||
Args:
|
||||
workspace_name: Name of the workspace
|
||||
session_name: Name of the session
|
||||
|
||||
Returns:
|
||||
Select statement returning peer_name, peer_configuration, and session_peer_configuration
|
||||
Select statement returning peer_name, peer_configuration, session_peer_configuration,
|
||||
and a boolean indicating if the peer is currently in the session
|
||||
"""
|
||||
stmt: Select[tuple[str, dict[str, Any], dict[str, Any]]] = (
|
||||
stmt: Select[tuple[str, dict[str, Any], dict[str, Any], bool]] = (
|
||||
select(
|
||||
models.Peer.name.label("peer_name"),
|
||||
models.Peer.configuration.label("peer_configuration"),
|
||||
models.SessionPeer.configuration.label("session_peer_configuration"),
|
||||
(models.SessionPeer.left_at.is_(None)).label("is_active"),
|
||||
)
|
||||
.join(models.SessionPeer, models.Peer.name == models.SessionPeer.peer_name)
|
||||
.where(models.SessionPeer.session_name == session_name)
|
||||
.where(models.Peer.workspace_name == workspace_name)
|
||||
.where(models.SessionPeer.workspace_name == workspace_name)
|
||||
.where(models.SessionPeer.left_at.is_(None)) # Only active peers
|
||||
)
|
||||
|
||||
return stmt
|
||||
|
|
|
|||
|
|
@ -118,7 +118,11 @@ async def get_peers_with_configuration(
|
|||
peers_with_configuration_result = await db_session.execute(configuration_query)
|
||||
peers_with_configuration_list = peers_with_configuration_result.all()
|
||||
return {
|
||||
row.peer_name: [row.peer_configuration, row.session_peer_configuration]
|
||||
row.peer_name: [
|
||||
row.peer_configuration,
|
||||
row.session_peer_configuration,
|
||||
row.is_active,
|
||||
]
|
||||
for row in peers_with_configuration_list
|
||||
}
|
||||
|
||||
|
|
@ -194,7 +198,10 @@ def get_effective_observe_me(
|
|||
Returns:
|
||||
True if observe_me is enabled, False otherwise
|
||||
"""
|
||||
configuration = peers_with_configuration[sender_name]
|
||||
# If the sender is not in peers_with_configuration, they left after sending a message.
|
||||
# We'll use the default behavior of observing the sender by instantiating the default
|
||||
# peer-level and session-level configs.
|
||||
configuration: list[Any] = peers_with_configuration.get(sender_name, [{}, {}])
|
||||
sender_session_peer_config = (
|
||||
schemas.SessionPeerConfig(**configuration[1]) if configuration[1] else None
|
||||
)
|
||||
|
|
@ -274,6 +281,10 @@ async def generate_queue_records(
|
|||
if peer_name == sender_name:
|
||||
continue
|
||||
|
||||
# If the observer peer has left the session, we don't need to enqueue a representation task for them.
|
||||
if not configuration[2]:
|
||||
continue
|
||||
|
||||
session_peer_config = (
|
||||
schemas.SessionPeerConfig(**configuration[1])
|
||||
if configuration[1]
|
||||
|
|
|
|||
|
|
@ -542,3 +542,771 @@ class TestEnqueueFunction:
|
|||
# For each expected payload, assert it is present in actual_payloads
|
||||
for expected in expected_payloads:
|
||||
assert expected in actual_payloads
|
||||
|
||||
# RACE CONDITION TESTS - Testing the new logic for peers that have left
|
||||
@pytest.mark.asyncio
|
||||
@patch("src.deriver.enqueue.tracked_db")
|
||||
async def test_sender_left_session_after_message_sent(
|
||||
self,
|
||||
mock_tracked_db: AsyncMock,
|
||||
db_session: AsyncSession,
|
||||
sample_data: tuple[Workspace, Peer],
|
||||
):
|
||||
"""Test that messages from senders who left the session still get processed with default config"""
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, sender_peer = sample_data
|
||||
|
||||
# Create an observer peer
|
||||
observer_peer = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
db_session.add(observer_peer)
|
||||
|
||||
# Create session with both peers
|
||||
test_session = await crud.get_or_create_session(
|
||||
db_session,
|
||||
schemas.SessionCreate(
|
||||
name=str(generate_nanoid()),
|
||||
peers={
|
||||
sender_peer.name: schemas.SessionPeerConfig(observe_me=True),
|
||||
observer_peer.name: schemas.SessionPeerConfig(observe_others=True),
|
||||
},
|
||||
),
|
||||
test_workspace.name,
|
||||
)
|
||||
await db_session.commit()
|
||||
|
||||
# Simulate sender leaving the session by setting left_at
|
||||
from datetime import datetime, timezone
|
||||
|
||||
session_peer_result = await db_session.execute(
|
||||
select(models.SessionPeer).where(
|
||||
models.SessionPeer.session_name == test_session.name,
|
||||
models.SessionPeer.peer_name == sender_peer.name,
|
||||
models.SessionPeer.workspace_name == test_workspace.name,
|
||||
)
|
||||
)
|
||||
session_peer = session_peer_result.scalar_one()
|
||||
session_peer.left_at = datetime.now(timezone.utc)
|
||||
await db_session.commit()
|
||||
|
||||
# Create message payload from the peer who left
|
||||
payload = self.create_sample_payload(
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=sender_peer.name,
|
||||
)
|
||||
|
||||
initial_count = await self.count_queue_items(db_session)
|
||||
await enqueue(payload)
|
||||
final_count = await self.count_queue_items(db_session)
|
||||
|
||||
# Should create 2 queue items:
|
||||
# 1 representation for sender (using default config since they left)
|
||||
# 1 representation for observer (still in session and observing others)
|
||||
assert final_count - initial_count == 2
|
||||
|
||||
result = await db_session.execute(
|
||||
select(QueueItem).where(QueueItem.session_id == test_session.id)
|
||||
)
|
||||
queue_items = result.scalars().all()
|
||||
|
||||
expected_payloads = [
|
||||
{
|
||||
"sender_name": sender_peer.name,
|
||||
"target_name": sender_peer.name,
|
||||
"task_type": "representation",
|
||||
},
|
||||
{
|
||||
"sender_name": sender_peer.name,
|
||||
"target_name": observer_peer.name,
|
||||
"task_type": "representation",
|
||||
},
|
||||
]
|
||||
actual_payloads = [
|
||||
{
|
||||
"sender_name": item.payload.get("sender_name"),
|
||||
"target_name": item.payload.get("target_name"),
|
||||
"task_type": item.payload.get("task_type"),
|
||||
}
|
||||
for item in queue_items
|
||||
]
|
||||
|
||||
assert len(actual_payloads) == len(expected_payloads)
|
||||
for expected in expected_payloads:
|
||||
assert expected in actual_payloads
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("src.deriver.enqueue.tracked_db")
|
||||
async def test_observer_left_session_no_queue_items_generated(
|
||||
self,
|
||||
mock_tracked_db: AsyncMock,
|
||||
db_session: AsyncSession,
|
||||
sample_data: tuple[Workspace, Peer],
|
||||
):
|
||||
"""Test that peers who left the session don't get representation tasks enqueued"""
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, sender_peer = sample_data
|
||||
|
||||
# Create observer peers - one will leave, one will stay
|
||||
observer_who_left = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
observer_who_stayed = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
db_session.add_all([observer_who_left, observer_who_stayed])
|
||||
|
||||
# Create session with all peers
|
||||
test_session = await crud.get_or_create_session(
|
||||
db_session,
|
||||
schemas.SessionCreate(
|
||||
name=str(generate_nanoid()),
|
||||
peers={
|
||||
sender_peer.name: schemas.SessionPeerConfig(observe_me=True),
|
||||
observer_who_left.name: schemas.SessionPeerConfig(
|
||||
observe_others=True
|
||||
),
|
||||
observer_who_stayed.name: schemas.SessionPeerConfig(
|
||||
observe_others=True
|
||||
),
|
||||
},
|
||||
),
|
||||
test_workspace.name,
|
||||
)
|
||||
await db_session.commit()
|
||||
|
||||
# Simulate one observer leaving the session
|
||||
from datetime import datetime, timezone
|
||||
|
||||
session_peer_result = await db_session.execute(
|
||||
select(models.SessionPeer).where(
|
||||
models.SessionPeer.session_name == test_session.name,
|
||||
models.SessionPeer.peer_name == observer_who_left.name,
|
||||
models.SessionPeer.workspace_name == test_workspace.name,
|
||||
)
|
||||
)
|
||||
session_peer = session_peer_result.scalar_one()
|
||||
session_peer.left_at = datetime.now(timezone.utc)
|
||||
await db_session.commit()
|
||||
|
||||
# Create message payload
|
||||
payload = self.create_sample_payload(
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=sender_peer.name,
|
||||
)
|
||||
|
||||
initial_count = await self.count_queue_items(db_session)
|
||||
await enqueue(payload)
|
||||
final_count = await self.count_queue_items(db_session)
|
||||
|
||||
# Should create 2 queue items:
|
||||
# 1 representation for sender
|
||||
# 1 representation for observer_who_stayed (observer_who_left should be skipped)
|
||||
assert final_count - initial_count == 2
|
||||
|
||||
result = await db_session.execute(
|
||||
select(QueueItem).where(QueueItem.session_id == test_session.id)
|
||||
)
|
||||
queue_items = result.scalars().all()
|
||||
|
||||
# Verify observer_who_left is NOT in the target names
|
||||
target_names = [item.payload.get("target_name") for item in queue_items]
|
||||
assert observer_who_left.name not in target_names
|
||||
assert observer_who_stayed.name in target_names
|
||||
assert sender_peer.name in target_names
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("src.deriver.enqueue.tracked_db")
|
||||
async def test_sender_not_in_peer_configuration_uses_defaults(
|
||||
self,
|
||||
mock_tracked_db: AsyncMock,
|
||||
db_session: AsyncSession,
|
||||
sample_data: tuple[Workspace, Peer],
|
||||
):
|
||||
"""Test get_effective_observe_me handles missing sender configuration gracefully"""
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, _existing_peer = sample_data
|
||||
|
||||
# Create observer peer
|
||||
observer_peer = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
db_session.add(observer_peer)
|
||||
|
||||
# Create session with only observer (sender not in peers_with_configuration)
|
||||
test_session = await crud.get_or_create_session(
|
||||
db_session,
|
||||
schemas.SessionCreate(
|
||||
name=str(generate_nanoid()),
|
||||
peers={
|
||||
observer_peer.name: schemas.SessionPeerConfig(observe_others=True),
|
||||
},
|
||||
),
|
||||
test_workspace.name,
|
||||
)
|
||||
await db_session.commit()
|
||||
|
||||
# Create message from peer NOT in the session configuration
|
||||
# This simulates the race condition where a peer left after sending
|
||||
unknown_sender = str(generate_nanoid())
|
||||
payload = self.create_sample_payload(
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=unknown_sender,
|
||||
)
|
||||
|
||||
initial_count = await self.count_queue_items(db_session)
|
||||
await enqueue(payload)
|
||||
final_count = await self.count_queue_items(db_session)
|
||||
|
||||
# Should create 2 queue items:
|
||||
# 1 representation for unknown sender (using default observe_me=True)
|
||||
# 1 representation for observer (observing others)
|
||||
assert final_count - initial_count == 2
|
||||
|
||||
result = await db_session.execute(
|
||||
select(QueueItem).where(QueueItem.session_id == test_session.id)
|
||||
)
|
||||
queue_items = result.scalars().all()
|
||||
|
||||
expected_payloads = [
|
||||
{
|
||||
"sender_name": unknown_sender,
|
||||
"target_name": unknown_sender,
|
||||
"task_type": "representation",
|
||||
},
|
||||
{
|
||||
"sender_name": unknown_sender,
|
||||
"target_name": observer_peer.name,
|
||||
"task_type": "representation",
|
||||
},
|
||||
]
|
||||
actual_payloads = [
|
||||
{
|
||||
"sender_name": item.payload.get("sender_name"),
|
||||
"target_name": item.payload.get("target_name"),
|
||||
"task_type": item.payload.get("task_type"),
|
||||
}
|
||||
for item in queue_items
|
||||
]
|
||||
|
||||
assert len(actual_payloads) == len(expected_payloads)
|
||||
for expected in expected_payloads:
|
||||
assert expected in actual_payloads
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("src.deriver.enqueue.tracked_db")
|
||||
async def test_mixed_active_inactive_peers_complex_scenario(
|
||||
self,
|
||||
mock_tracked_db: AsyncMock,
|
||||
db_session: AsyncSession,
|
||||
sample_data: tuple[Workspace, Peer],
|
||||
):
|
||||
"""Test complex scenario with mix of active/inactive peers and different configurations"""
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, sender_peer = sample_data
|
||||
|
||||
# Create multiple peers with different roles
|
||||
active_observer = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
inactive_observer = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
active_non_observer = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
inactive_non_observer = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
|
||||
db_session.add_all(
|
||||
[
|
||||
active_observer,
|
||||
inactive_observer,
|
||||
active_non_observer,
|
||||
inactive_non_observer,
|
||||
]
|
||||
)
|
||||
|
||||
# Create session with all peers having different configurations
|
||||
test_session = await crud.get_or_create_session(
|
||||
db_session,
|
||||
schemas.SessionCreate(
|
||||
name=str(generate_nanoid()),
|
||||
peers={
|
||||
sender_peer.name: schemas.SessionPeerConfig(observe_me=True),
|
||||
active_observer.name: schemas.SessionPeerConfig(
|
||||
observe_others=True
|
||||
),
|
||||
inactive_observer.name: schemas.SessionPeerConfig(
|
||||
observe_others=True
|
||||
),
|
||||
active_non_observer.name: schemas.SessionPeerConfig(
|
||||
observe_others=False
|
||||
),
|
||||
inactive_non_observer.name: schemas.SessionPeerConfig(
|
||||
observe_others=False
|
||||
),
|
||||
},
|
||||
),
|
||||
test_workspace.name,
|
||||
)
|
||||
await db_session.commit()
|
||||
|
||||
# Mark some peers as having left the session
|
||||
from datetime import datetime, timezone
|
||||
|
||||
for peer_name in [inactive_observer.name, inactive_non_observer.name]:
|
||||
session_peer_result = await db_session.execute(
|
||||
select(models.SessionPeer).where(
|
||||
models.SessionPeer.session_name == test_session.name,
|
||||
models.SessionPeer.peer_name == peer_name,
|
||||
models.SessionPeer.workspace_name == test_workspace.name,
|
||||
)
|
||||
)
|
||||
session_peer = session_peer_result.scalar_one()
|
||||
session_peer.left_at = datetime.now(timezone.utc)
|
||||
await db_session.commit()
|
||||
|
||||
# Create message payload from sender
|
||||
payload = self.create_sample_payload(
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=sender_peer.name,
|
||||
)
|
||||
|
||||
initial_count = await self.count_queue_items(db_session)
|
||||
await enqueue(payload)
|
||||
final_count = await self.count_queue_items(db_session)
|
||||
|
||||
# Should create 2 queue items:
|
||||
# 1 representation for sender (observe_me=True)
|
||||
# 1 representation for active_observer (observe_others=True and still active)
|
||||
# inactive_observer should be skipped (left session)
|
||||
# active_non_observer should be skipped (observe_others=False)
|
||||
# inactive_non_observer should be skipped (left session)
|
||||
assert final_count - initial_count == 2
|
||||
|
||||
result = await db_session.execute(
|
||||
select(QueueItem).where(QueueItem.session_id == test_session.id)
|
||||
)
|
||||
queue_items = result.scalars().all()
|
||||
|
||||
expected_payloads = [
|
||||
{
|
||||
"sender_name": sender_peer.name,
|
||||
"target_name": sender_peer.name,
|
||||
"task_type": "representation",
|
||||
},
|
||||
{
|
||||
"sender_name": sender_peer.name,
|
||||
"target_name": active_observer.name,
|
||||
"task_type": "representation",
|
||||
},
|
||||
]
|
||||
actual_payloads = [
|
||||
{
|
||||
"sender_name": item.payload.get("sender_name"),
|
||||
"target_name": item.payload.get("target_name"),
|
||||
"task_type": item.payload.get("task_type"),
|
||||
}
|
||||
for item in queue_items
|
||||
]
|
||||
|
||||
assert len(actual_payloads) == len(expected_payloads)
|
||||
for expected in expected_payloads:
|
||||
assert expected in actual_payloads
|
||||
|
||||
# Verify inactive peers are not in target names
|
||||
target_names = [item.payload.get("target_name") for item in queue_items]
|
||||
assert inactive_observer.name not in target_names
|
||||
assert inactive_non_observer.name not in target_names
|
||||
assert active_non_observer.name not in target_names
|
||||
|
||||
|
||||
class TestGetEffectiveObserveMeFunction:
|
||||
"""Unit tests for the get_effective_observe_me function specifically testing race condition handling"""
|
||||
|
||||
def test_sender_missing_from_configuration_uses_default(self):
|
||||
"""Test that missing sender uses default observe_me=True"""
|
||||
from src.deriver.enqueue import get_effective_observe_me
|
||||
|
||||
# Empty peer configuration dict simulates sender who left after sending message
|
||||
peers_with_configuration: dict[str, list[dict[str, Any]]] = {}
|
||||
|
||||
result = get_effective_observe_me("missing_sender", peers_with_configuration)
|
||||
|
||||
# Should use default PeerConfig() which has observe_me=True
|
||||
assert result is True
|
||||
|
||||
def test_sender_with_empty_configurations_uses_default(self):
|
||||
"""Test that sender with empty peer and session configs uses default"""
|
||||
from src.deriver.enqueue import get_effective_observe_me
|
||||
|
||||
# Sender present but with empty configurations
|
||||
peers_with_configuration: dict[str, list[dict[str, Any]]] = {
|
||||
"sender": [{}, {}] # Empty peer config, empty session config
|
||||
}
|
||||
|
||||
result = get_effective_observe_me("sender", peers_with_configuration)
|
||||
|
||||
# Should use default PeerConfig() which has observe_me=True
|
||||
assert result is True
|
||||
|
||||
def test_sender_with_peer_config_observe_me_false(self):
|
||||
"""Test that peer config observe_me=False is respected"""
|
||||
from src.deriver.enqueue import get_effective_observe_me
|
||||
|
||||
peers_with_configuration = {
|
||||
"sender": [{"observe_me": False}, {}] # Peer config with observe_me=False
|
||||
}
|
||||
|
||||
result = get_effective_observe_me("sender", peers_with_configuration)
|
||||
|
||||
assert result is False
|
||||
|
||||
def test_session_config_overrides_peer_config(self):
|
||||
"""Test that session peer config takes precedence over peer config"""
|
||||
from src.deriver.enqueue import get_effective_observe_me
|
||||
|
||||
peers_with_configuration = {
|
||||
"sender": [
|
||||
{"observe_me": True}, # Peer config says True
|
||||
{"observe_me": False}, # Session config says False - should win
|
||||
]
|
||||
}
|
||||
|
||||
result = get_effective_observe_me("sender", peers_with_configuration)
|
||||
|
||||
assert result is False
|
||||
|
||||
def test_session_config_none_falls_back_to_peer_config(self):
|
||||
"""Test that session config with None observe_me falls back to peer config"""
|
||||
from src.deriver.enqueue import get_effective_observe_me
|
||||
|
||||
peers_with_configuration = {
|
||||
"sender": [
|
||||
{"observe_me": False}, # Peer config says False
|
||||
{"observe_me": None}, # Session config is None - should fall back
|
||||
]
|
||||
}
|
||||
|
||||
result = get_effective_observe_me("sender", peers_with_configuration)
|
||||
|
||||
assert result is False
|
||||
|
||||
def test_mixed_configurations_with_active_status(self):
|
||||
"""Test various configuration combinations with active status"""
|
||||
from src.deriver.enqueue import get_effective_observe_me
|
||||
|
||||
# Test cases: (peer_config, session_config, expected_result)
|
||||
test_cases: list[tuple[dict[str, Any] | None, dict[str, Any] | None, bool]] = [
|
||||
# Default case - missing sender
|
||||
(None, None, True),
|
||||
# Empty configs
|
||||
({}, {}, True),
|
||||
# Peer config only
|
||||
({"observe_me": False}, {}, False),
|
||||
({"observe_me": True}, {}, True),
|
||||
# Session config overrides
|
||||
({"observe_me": True}, {"observe_me": False}, False),
|
||||
({"observe_me": False}, {"observe_me": True}, True),
|
||||
# Session config None falls back to peer
|
||||
({"observe_me": False}, {"observe_me": None}, False),
|
||||
({"observe_me": True}, {"observe_me": None}, True),
|
||||
]
|
||||
|
||||
for i, (peer_config, session_config, expected) in enumerate(test_cases):
|
||||
if peer_config is None:
|
||||
# Test missing sender
|
||||
peers_with_configuration = {}
|
||||
sender_name = "missing_sender"
|
||||
else:
|
||||
peers_with_configuration = {
|
||||
f"sender_{i}": [peer_config or {}, session_config or {}]
|
||||
}
|
||||
sender_name = f"sender_{i}"
|
||||
|
||||
result = get_effective_observe_me(sender_name, peers_with_configuration)
|
||||
assert result == expected, (
|
||||
f"Test case {i} failed: peer_config={peer_config}, session_config={session_config}, expected={expected}, got={result}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestAdvancedEnqueueEdgeCases:
|
||||
"""Test advanced edge cases for the enqueue system with race conditions"""
|
||||
|
||||
# Helper methods
|
||||
def create_sample_payload(
|
||||
self,
|
||||
workspace_name: str = "test_workspace",
|
||||
session_name: str | None = "test_session",
|
||||
peer_name: str = "test_peer",
|
||||
count: int = 1,
|
||||
):
|
||||
"""Create sample payload for testing"""
|
||||
return [
|
||||
{
|
||||
"workspace_name": workspace_name,
|
||||
"session_name": session_name,
|
||||
"message_id": i + 1,
|
||||
"content": f"Test message {i}",
|
||||
"metadata": {"test": f"value_{i}"},
|
||||
"peer_name": peer_name,
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
}
|
||||
for i in range(count)
|
||||
]
|
||||
|
||||
async def count_queue_items(self, db_session: AsyncSession):
|
||||
"""Helper to count queue items in database"""
|
||||
result = await db_session.execute(select(QueueItem))
|
||||
return len(result.scalars().all())
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("src.deriver.enqueue.tracked_db")
|
||||
async def test_edge_case_all_peers_left_except_sender(
|
||||
self,
|
||||
mock_tracked_db: AsyncMock,
|
||||
db_session: AsyncSession,
|
||||
sample_data: tuple[Workspace, Peer],
|
||||
):
|
||||
"""Test edge case where all observer peers have left the session"""
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, sender_peer = sample_data
|
||||
|
||||
# Create multiple observer peers
|
||||
observer1 = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
observer2 = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
db_session.add_all([observer1, observer2])
|
||||
|
||||
# Create session with all peers
|
||||
test_session = await crud.get_or_create_session(
|
||||
db_session,
|
||||
schemas.SessionCreate(
|
||||
name=str(generate_nanoid()),
|
||||
peers={
|
||||
sender_peer.name: schemas.SessionPeerConfig(observe_me=True),
|
||||
observer1.name: schemas.SessionPeerConfig(observe_others=True),
|
||||
observer2.name: schemas.SessionPeerConfig(observe_others=True),
|
||||
},
|
||||
),
|
||||
test_workspace.name,
|
||||
)
|
||||
await db_session.commit()
|
||||
|
||||
# Mark all observers as having left
|
||||
from datetime import datetime, timezone
|
||||
|
||||
for peer_name in [observer1.name, observer2.name]:
|
||||
session_peer_result = await db_session.execute(
|
||||
select(models.SessionPeer).where(
|
||||
models.SessionPeer.session_name == test_session.name,
|
||||
models.SessionPeer.peer_name == peer_name,
|
||||
models.SessionPeer.workspace_name == test_workspace.name,
|
||||
)
|
||||
)
|
||||
session_peer = session_peer_result.scalar_one()
|
||||
session_peer.left_at = datetime.now(timezone.utc)
|
||||
await db_session.commit()
|
||||
|
||||
# Create message payload
|
||||
payload = self.create_sample_payload(
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=sender_peer.name,
|
||||
)
|
||||
|
||||
initial_count = await self.count_queue_items(db_session)
|
||||
await enqueue(payload)
|
||||
final_count = await self.count_queue_items(db_session)
|
||||
|
||||
# Should create only 1 queue item: representation for sender only
|
||||
assert final_count - initial_count == 1
|
||||
|
||||
result = await db_session.execute(
|
||||
select(QueueItem).where(QueueItem.session_id == test_session.id)
|
||||
)
|
||||
queue_items = result.scalars().all()
|
||||
|
||||
assert len(queue_items) == 1
|
||||
assert queue_items[0].payload["sender_name"] == sender_peer.name
|
||||
assert queue_items[0].payload["target_name"] == sender_peer.name
|
||||
assert queue_items[0].payload["task_type"] == "representation"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("src.deriver.enqueue.tracked_db")
|
||||
async def test_edge_case_sender_and_observer_both_left_different_times(
|
||||
self,
|
||||
mock_tracked_db: AsyncMock,
|
||||
db_session: AsyncSession,
|
||||
sample_data: tuple[Workspace, Peer],
|
||||
):
|
||||
"""Test race condition where both sender and observer left at different times"""
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, sender_peer = sample_data
|
||||
|
||||
observer_peer = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
db_session.add(observer_peer)
|
||||
|
||||
# Create session
|
||||
test_session = await crud.get_or_create_session(
|
||||
db_session,
|
||||
schemas.SessionCreate(
|
||||
name=str(generate_nanoid()),
|
||||
peers={
|
||||
sender_peer.name: schemas.SessionPeerConfig(observe_me=True),
|
||||
observer_peer.name: schemas.SessionPeerConfig(observe_others=True),
|
||||
},
|
||||
),
|
||||
test_workspace.name,
|
||||
)
|
||||
await db_session.commit()
|
||||
|
||||
# Mark both as having left (observer left first, then sender)
|
||||
from datetime import datetime, timezone
|
||||
|
||||
base_time = datetime.now(timezone.utc)
|
||||
|
||||
# Observer left first
|
||||
observer_session_peer_result = await db_session.execute(
|
||||
select(models.SessionPeer).where(
|
||||
models.SessionPeer.session_name == test_session.name,
|
||||
models.SessionPeer.peer_name == observer_peer.name,
|
||||
models.SessionPeer.workspace_name == test_workspace.name,
|
||||
)
|
||||
)
|
||||
observer_session_peer = observer_session_peer_result.scalar_one()
|
||||
observer_session_peer.left_at = base_time
|
||||
|
||||
# Sender left later
|
||||
sender_session_peer_result = await db_session.execute(
|
||||
select(models.SessionPeer).where(
|
||||
models.SessionPeer.session_name == test_session.name,
|
||||
models.SessionPeer.peer_name == sender_peer.name,
|
||||
models.SessionPeer.workspace_name == test_workspace.name,
|
||||
)
|
||||
)
|
||||
sender_session_peer = sender_session_peer_result.scalar_one()
|
||||
sender_session_peer.left_at = base_time
|
||||
|
||||
await db_session.commit()
|
||||
|
||||
# Create message payload from sender who left
|
||||
payload = self.create_sample_payload(
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=sender_peer.name,
|
||||
)
|
||||
|
||||
initial_count = await self.count_queue_items(db_session)
|
||||
await enqueue(payload)
|
||||
final_count = await self.count_queue_items(db_session)
|
||||
|
||||
# Should create only 1 queue item: representation for sender (using default config)
|
||||
# Observer should be skipped because they left the session
|
||||
assert final_count - initial_count == 1
|
||||
|
||||
result = await db_session.execute(
|
||||
select(QueueItem).where(QueueItem.session_id == test_session.id)
|
||||
)
|
||||
queue_items = result.scalars().all()
|
||||
|
||||
assert len(queue_items) == 1
|
||||
assert queue_items[0].payload["sender_name"] == sender_peer.name
|
||||
assert queue_items[0].payload["target_name"] == sender_peer.name
|
||||
assert queue_items[0].payload["task_type"] == "representation"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("src.deriver.enqueue.tracked_db")
|
||||
async def test_edge_case_message_from_never_joined_peer(
|
||||
self,
|
||||
mock_tracked_db: AsyncMock,
|
||||
db_session: AsyncSession,
|
||||
sample_data: tuple[Workspace, Peer],
|
||||
):
|
||||
"""Test handling message from peer who was never in the session"""
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, _existing_peer = sample_data
|
||||
|
||||
observer_peer = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
)
|
||||
db_session.add(observer_peer)
|
||||
|
||||
# Create session with only observer peer
|
||||
test_session = await crud.get_or_create_session(
|
||||
db_session,
|
||||
schemas.SessionCreate(
|
||||
name=str(generate_nanoid()),
|
||||
peers={
|
||||
observer_peer.name: schemas.SessionPeerConfig(observe_others=True),
|
||||
},
|
||||
),
|
||||
test_workspace.name,
|
||||
)
|
||||
await db_session.commit()
|
||||
|
||||
# Create message from peer who was NEVER in the session
|
||||
never_joined_peer = str(generate_nanoid())
|
||||
payload = self.create_sample_payload(
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=never_joined_peer,
|
||||
)
|
||||
|
||||
initial_count = await self.count_queue_items(db_session)
|
||||
await enqueue(payload)
|
||||
final_count = await self.count_queue_items(db_session)
|
||||
|
||||
# Should create 2 queue items:
|
||||
# 1 for never_joined_peer (using default config)
|
||||
# 1 for observer (observe_others=True)
|
||||
assert final_count - initial_count == 2
|
||||
|
||||
result = await db_session.execute(
|
||||
select(QueueItem).where(QueueItem.session_id == test_session.id)
|
||||
)
|
||||
queue_items = result.scalars().all()
|
||||
|
||||
expected_payloads = [
|
||||
{
|
||||
"sender_name": never_joined_peer,
|
||||
"target_name": never_joined_peer,
|
||||
"task_type": "representation",
|
||||
},
|
||||
{
|
||||
"sender_name": never_joined_peer,
|
||||
"target_name": observer_peer.name,
|
||||
"task_type": "representation",
|
||||
},
|
||||
]
|
||||
actual_payloads = [
|
||||
{
|
||||
"sender_name": item.payload.get("sender_name"),
|
||||
"target_name": item.payload.get("target_name"),
|
||||
"task_type": item.payload.get("task_type"),
|
||||
}
|
||||
for item in queue_items
|
||||
]
|
||||
|
||||
assert len(actual_payloads) == len(expected_payloads)
|
||||
for expected in expected_payloads:
|
||||
assert expected in actual_payloads
|
||||
|
|
|
|||
Loading…
Reference in New Issue