Use savepoints to prevent race condition in get_or_create chain (#399)

* fix: Use savepoints to prevent race condition in get_or_create chain

Replace commit()/rollback() with begin_nested() savepoints in
get_or_create_workspace, get_or_create_peers, and get_or_create_session
so that an IntegrityError rollback in a nested call doesn't undo flushed
work from the caller. Moves transaction commit responsibility to the
outermost caller and defers cache operations to post-commit callbacks
on GetOrCreateResult.

Closes DEV-1321

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* fix: Add missing commit and post_commit in chat endpoint

The chat() endpoint in routers/peers.py called get_or_create_peers()
but discarded the result without committing or invoking post_commit().
This meant new peers created lazily via SDK chat calls were never
persisted, and cache invalidation was skipped.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Vineeth Voruganti 2026-02-23 23:28:04 -05:00 committed by GitHub
parent 3b1e964372
commit 5a42f9b3cd
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
7 changed files with 91 additions and 57 deletions

View File

@ -113,29 +113,31 @@ async def get_or_create_peers(
for p in peers_to_create
]
try:
db.add_all(new_peers)
await db.commit()
async with db.begin_nested():
db.add_all(new_peers)
except IntegrityError:
await db.rollback()
if _retry:
raise ConflictException(
f"Unable to create or get peers: {peer_names}"
) from None
return await get_or_create_peers(db, workspace_name, peers, _retry=True)
# Only invalidate cache for changed/new peers - read-through pattern
for peer_obj in changed_peers + new_peers:
cache_key = peer_cache_key(workspace_name, peer_obj.name)
await safe_cache_delete(cache_key)
logger.debug(
"Peer %s cache invalidated in workspace %s (changed or new)",
peer_obj.name,
workspace_name,
)
# Capture peer names eagerly so the closure holds plain strings, not ORM objects
_cache_keys_to_invalidate = [
peer_cache_key(workspace_name, p.name) for p in changed_peers + new_peers
]
async def _invalidate_peer_cache():
for cache_key in _cache_keys_to_invalidate:
await safe_cache_delete(cache_key)
# Return combined list of existing and new peers
# created=True if any new peers were created
return GetOrCreateResult(existing_peers + new_peers, created=len(new_peers) > 0)
return GetOrCreateResult(
existing_peers + new_peers,
created=len(new_peers) > 0,
on_commit=_invalidate_peer_cache if _cache_keys_to_invalidate else None,
)
@cache(
@ -237,11 +239,10 @@ async def update_peer(
ValidationException: If the update data is invalid
ConflictException: If the update violates a unique constraint
"""
honcho_peer = (
await get_or_create_peers(
db, workspace_name, [schemas.PeerCreate(name=peer_name)]
)
).resource[0]
peers_result = await get_or_create_peers(
db, workspace_name, [schemas.PeerCreate(name=peer_name)]
)
honcho_peer = peers_result.resource[0]
needs_update = False
@ -258,6 +259,8 @@ async def update_peer(
# Early exit if unchanged
if not needs_update:
await db.commit()
await peers_result.post_commit()
logger.debug(
"Peer %s unchanged in workspace %s, skipping update",
peer_name,
@ -267,6 +270,7 @@ async def update_peer(
await db.commit()
await db.refresh(honcho_peer)
await peers_result.post_commit()
cache_key = peer_cache_key(workspace_name, honcho_peer.name)
await safe_cache_delete(cache_key)

View File

@ -69,7 +69,9 @@ async def set_peer_card(
"""
# Ensure the peer exists (get-or-create)
await get_or_create_peers(db, workspace_name, [schemas.PeerCreate(name=observer)])
peers_result = await get_or_create_peers(
db, workspace_name, [schemas.PeerCreate(name=observer)]
)
stmt = (
update(models.Peer)
@ -91,6 +93,7 @@ async def set_peer_card(
f"Peer {observer} not found in workspace {workspace_name}"
)
await db.commit()
await peers_result.post_commit()
# Invalidate cache - read-through pattern
cache_key = peer_cache_key(workspace_name, observer)

View File

@ -175,6 +175,8 @@ async def get_or_create_session(
# Track if we need to update cache and if session was created
needs_cache_update = False
created = False
ws_result = None
peers_result = None
# Check if session already exists
if honcho_session is None:
@ -185,7 +187,7 @@ async def get_or_create_session(
raise ObserverException(session.name, observer_count)
# Get or create workspace to ensure it exists
await get_or_create_workspace(
ws_result = await get_or_create_workspace(
db,
schemas.WorkspaceCreate(name=workspace_name),
)
@ -200,14 +202,12 @@ async def get_or_create_session(
else {},
)
try:
db.add(honcho_session)
# Flush to ensure session exists in DB before adding peers and set flag to warm cache
await db.flush()
async with db.begin_nested():
db.add(honcho_session)
needs_cache_update = True
created = True
except IntegrityError:
await db.rollback()
logger.debug(
"Race condition detected for session: %s, retrying get", session.name
)
@ -235,7 +235,7 @@ async def get_or_create_session(
# Add all peers to session
if session.peer_names:
await get_or_create_peers(
peers_result = await get_or_create_peers(
db,
workspace_name=workspace_name,
peers=[
@ -252,6 +252,12 @@ async def get_or_create_session(
await db.commit()
await db.refresh(honcho_session)
# Run deferred cache operations from workspace/peer creation
if ws_result is not None:
await ws_result.post_commit()
if peers_result is not None:
await peers_result.post_commit()
# Only update cache if session data changed or was newly created
if needs_cache_update:
cache_key = session_cache_key(workspace_name, session.name)
@ -882,7 +888,7 @@ async def set_peers_for_session(
result = await db.execute(update_stmt)
# Get or create peers
await get_or_create_peers(
peers_result = await get_or_create_peers(
db,
workspace_name=workspace_name,
peers=[schemas.PeerCreate(name=peer_name) for peer_name in peer_names],
@ -897,6 +903,7 @@ async def set_peers_for_session(
)
await db.commit()
await peers_result.post_commit()
return peers

View File

@ -119,28 +119,32 @@ async def get_or_create_workspace(
configuration=workspace.configuration.model_dump(exclude_none=True),
)
try:
db.add(honcho_workspace)
await db.commit()
await db.refresh(honcho_workspace)
async with db.begin_nested():
db.add(honcho_workspace)
logger.debug("Workspace created successfully: %s", workspace.name)
cache_key = workspace_cache_key(workspace.name)
await safe_cache_set(
cache_key,
{
"id": honcho_workspace.id,
"name": honcho_workspace.name,
"h_metadata": honcho_workspace.h_metadata,
"internal_metadata": honcho_workspace.internal_metadata,
"configuration": honcho_workspace.configuration,
"created_at": honcho_workspace.created_at,
},
expire=settings.CACHE.DEFAULT_TTL_SECONDS,
# Capture cache data eagerly so the closure holds a plain dict, not the ORM object
_cache_key = workspace_cache_key(workspace.name)
_cache_data = {
"id": honcho_workspace.id,
"name": honcho_workspace.name,
"h_metadata": honcho_workspace.h_metadata,
"internal_metadata": honcho_workspace.internal_metadata,
"configuration": honcho_workspace.configuration,
"created_at": honcho_workspace.created_at,
}
async def _warm_workspace_cache():
await safe_cache_set(
_cache_key,
_cache_data,
expire=settings.CACHE.DEFAULT_TTL_SECONDS,
)
return GetOrCreateResult(
honcho_workspace, created=True, on_commit=_warm_workspace_cache
)
return GetOrCreateResult(honcho_workspace, created=True)
except IntegrityError:
await db.rollback()
if _retry:
raise ConflictException(
f"Unable to create or get workspace: {workspace.name}"
@ -208,16 +212,14 @@ async def update_workspace(
Returns:
The updated workspace
"""
honcho_workspace: models.Workspace = (
await get_or_create_workspace(
db,
schemas.WorkspaceCreate(
name=workspace_name,
metadata=workspace.metadata
or {}, # Provide empty dict if metadata is None
),
)
).resource
ws_result = await get_or_create_workspace(
db,
schemas.WorkspaceCreate(
name=workspace_name,
metadata=workspace.metadata or {}, # Provide empty dict if metadata is None
),
)
honcho_workspace: models.Workspace = ws_result.resource
# Track if anything changed
needs_update = False
@ -242,11 +244,14 @@ async def update_workspace(
# Early exit if unchanged
if not needs_update:
await db.commit()
await ws_result.post_commit()
logger.debug("Workspace %s unchanged, skipping update", workspace_name)
return honcho_workspace
await db.commit()
await db.refresh(honcho_workspace)
await ws_result.post_commit()
# Only invalidate if we actually updated
cache_key = workspace_cache_key(workspace_name)

View File

@ -82,6 +82,8 @@ async def get_or_create_peer(
result = await crud.get_or_create_peers(
db, workspace_name=workspace_id, peers=[peer]
)
await db.commit()
await result.post_commit()
response.status_code = 201 if result.created else 200
return result.resource[0]
@ -167,11 +169,13 @@ async def chat(
"""
# Get or create the peer to ensure it exists
async with tracked_db("peers.chat.get_or_create_peer") as peer_db:
await crud.get_or_create_peers(
peers_result = await crud.get_or_create_peers(
peer_db,
workspace_name=workspace_id,
peers=[schemas.PeerCreate(name=peer_id)],
)
await peer_db.commit()
await peers_result.post_commit()
if options.stream:
# Stream the response using Server-Sent Events

View File

@ -50,6 +50,8 @@ async def get_or_create_workspace(
workspace.name = jwt_params.w
result = await crud.get_or_create_workspace(db, workspace=workspace)
await db.commit()
await result.post_commit()
response.status_code = 201 if result.created else 200
return result.resource

View File

@ -1,5 +1,7 @@
from collections.abc import Awaitable, Callable
from contextvars import ContextVar
from typing import Generic, Literal, NamedTuple, TypeVar
from dataclasses import dataclass, field
from typing import Generic, Literal, TypeVar
T = TypeVar("T")
@ -18,11 +20,18 @@ def get_current_iteration() -> int:
return _current_iteration.get()
class GetOrCreateResult(NamedTuple, Generic[T]):
@dataclass
class GetOrCreateResult(Generic[T]):
"""Result of a get_or_create operation indicating whether the resource was created."""
resource: T
created: bool
on_commit: Callable[[], Awaitable[None]] | None = field(default=None, repr=False)
async def post_commit(self) -> None:
"""Run deferred cache operations after the transaction is committed."""
if self.on_commit is not None:
await self.on_commit()
SupportedProviders = Literal["anthropic", "openai", "google", "groq", "custom", "vllm"]