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:
parent
3b1e964372
commit
5a42f9b3cd
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Reference in New Issue