Fix concurrent thread metadata merges (#4489)

This commit is contained in:
Daoyuan Li 2026-07-27 07:18:02 -07:00 committed by GitHub
parent 5ddb678bc3
commit 5ce3cecf2a
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 31 additions and 5 deletions

View File

@ -6,7 +6,7 @@ import logging
from datetime import UTC, datetime
from typing import Any
from sqlalchemy import case, select, update
from sqlalchemy import case, select, text, update
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from sqlalchemy.orm.attributes import flag_modified
@ -204,16 +204,27 @@ class ThreadMetaRepository(ThreadMetaStore):
) -> None:
"""Merge ``metadata`` into ``metadata_json``.
Read-modify-write inside a single session/transaction so concurrent
callers see consistent state. No-op if the row does not exist or
the user_id check fails.
The row is locked before the read-modify-write merge so concurrent
callers cannot replace each other's keys. SQLite acquires its write
transaction before reading; databases with row-level locking use
``SELECT ... FOR UPDATE``. No-op if the row does not exist or the
user_id check fails.
``touch`` refreshes ``updated_at`` (default); pass ``touch=False`` to
preserve recency ordering for metadata-only changes such as pin/unpin.
"""
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_metadata")
async with self._sf() as session:
row = await session.get(ThreadMetaRow, thread_id)
if session.get_bind().dialect.name == "sqlite":
# A deferred SQLite transaction does not reserve the writer
# until the UPDATE, which is too late for a read-modify-write
# merge. BEGIN IMMEDIATE serializes writers before the read,
# including writers in other processes using the same file.
await session.execute(text("BEGIN IMMEDIATE"))
row = await session.get(ThreadMetaRow, thread_id)
else:
result = await session.execute(select(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id).with_for_update())
row = result.scalar_one_or_none()
if row is None:
return
if resolved_user_id is not None and row.user_id != resolved_user_id:

View File

@ -1,5 +1,6 @@
"""Tests for ThreadMetaRepository (SQLAlchemy-backed)."""
import asyncio
import logging
import pytest
@ -160,6 +161,20 @@ class TestThreadMetaRepository:
assert record["metadata"] == {"a": 1, THREAD_PINNED_METADATA_KEY: True}
assert record["updated_at"] == original
@pytest.mark.anyio
async def test_concurrent_metadata_updates_preserve_disjoint_keys(self, repo):
for index in range(10):
thread_id = f"concurrent-{index}"
await repo.create(thread_id, metadata={"base": index}, user_id=None)
await asyncio.gather(
repo.update_metadata(thread_id, {"left": index}, user_id=None),
repo.update_metadata(thread_id, {"right": index}, user_id=None),
)
record = await repo.get(thread_id, user_id=None)
assert record["metadata"] == {"base": index, "left": index, "right": index}
@pytest.mark.anyio
async def test_search_orders_pinned_threads_before_newer_unpinned_threads(self, repo):
await repo.create("older-pinned", metadata={THREAD_PINNED_METADATA_KEY: True})