mirror of
https://github.com/crewAIInc/crewAI.git
synced 2026-09-20 18:13:49 +00:00
fix(flows): order MongoDB state writes atomically
This commit is contained in:
@@ -154,7 +154,7 @@ class MongoDbFlowPersistence(FlowPersistence):
|
||||
f"state_data must be either a Pydantic BaseModel or dict, got {type(state_data)}"
|
||||
)
|
||||
|
||||
def _next_sequence(self, name: str) -> int:
|
||||
def _next_sequence(self, name: str, session: Any) -> int:
|
||||
"""Return the next value of a server-assigned monotonic counter.
|
||||
|
||||
MongoDB has no autoincrement, so this atomically ``$inc`` a per-name
|
||||
@@ -169,6 +169,7 @@ class MongoDbFlowPersistence(FlowPersistence):
|
||||
{"$inc": {"seq": 1}},
|
||||
upsert=True,
|
||||
return_document=ReturnDocument.AFTER,
|
||||
session=session,
|
||||
)
|
||||
return int(doc["seq"])
|
||||
|
||||
@@ -185,15 +186,24 @@ class MongoDbFlowPersistence(FlowPersistence):
|
||||
ordering on ``seq`` rather than the client-generated ObjectId ``_id``.
|
||||
"""
|
||||
state_dict = self._to_state_dict(state_data)
|
||||
self._db_ready()[self.states_collection].insert_one(
|
||||
{
|
||||
"flow_uuid": flow_uuid,
|
||||
"method_name": method_name,
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"state_json": json.dumps(state_dict, default=_json_default),
|
||||
"seq": self._next_sequence(self.states_collection),
|
||||
}
|
||||
)
|
||||
db = self._db_ready()
|
||||
|
||||
def write_state(session: Any) -> None:
|
||||
db[self.states_collection].insert_one(
|
||||
{
|
||||
"flow_uuid": flow_uuid,
|
||||
"method_name": method_name,
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"state_json": json.dumps(state_dict, default=_json_default),
|
||||
"seq": self._next_sequence(self.states_collection, session),
|
||||
},
|
||||
session=session,
|
||||
)
|
||||
|
||||
if self._client is None:
|
||||
raise RuntimeError("MongoDB client was not initialized.")
|
||||
with self._client.start_session() as session:
|
||||
session.with_transaction(write_state)
|
||||
|
||||
def load_state(self, flow_uuid: str) -> dict[str, Any] | None:
|
||||
"""Load the most recent state for a given flow UUID."""
|
||||
|
||||
@@ -36,7 +36,7 @@ class _FakeCollection:
|
||||
def create_index(self, *args: Any, **kwargs: Any) -> None:
|
||||
pass
|
||||
|
||||
def insert_one(self, doc: dict[str, Any]) -> None:
|
||||
def insert_one(self, doc: dict[str, Any], session: Any = None) -> None:
|
||||
self.docs.append(dict(doc))
|
||||
|
||||
def find_one(
|
||||
@@ -55,9 +55,12 @@ class _FakeCollection:
|
||||
update: dict[str, Any],
|
||||
upsert: bool = False,
|
||||
return_document: Any = None,
|
||||
session: Any = None,
|
||||
) -> dict[str, Any] | None:
|
||||
row = next((d for d in self.docs if self._match(d, flt)), None)
|
||||
if row is None and upsert:
|
||||
if row is None:
|
||||
if not upsert:
|
||||
return None
|
||||
row = dict(flt)
|
||||
self.docs.append(row)
|
||||
for key, delta in update.get("$inc", {}).items():
|
||||
@@ -89,16 +92,35 @@ class _FakeDatabase:
|
||||
return self.collections.setdefault(name, _FakeCollection())
|
||||
|
||||
|
||||
class _FakeSession:
|
||||
def __init__(self, client: _FakeClient) -> None:
|
||||
self.client = client
|
||||
|
||||
def __enter__(self) -> _FakeSession:
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: Any) -> None:
|
||||
return None
|
||||
|
||||
def with_transaction(self, callback: Any) -> None:
|
||||
self.client.transactions_started += 1
|
||||
callback(self)
|
||||
|
||||
|
||||
class _FakeClient:
|
||||
def __init__(self, conn: str) -> None:
|
||||
self.conn = conn
|
||||
self.db_names: list[str] = []
|
||||
self._db = _FakeDatabase()
|
||||
self.transactions_started = 0
|
||||
|
||||
def __getitem__(self, name: str) -> _FakeDatabase:
|
||||
self.db_names.append(name)
|
||||
return self._db
|
||||
|
||||
def start_session(self) -> _FakeSession:
|
||||
return _FakeSession(self)
|
||||
|
||||
|
||||
def _patch_client(monkeypatch: pytest.MonkeyPatch) -> dict[str, Any]:
|
||||
"""Patch ``pymongo.MongoClient`` to build in-memory fakes.
|
||||
@@ -155,6 +177,17 @@ def test_save_state_tags_incrementing_seq(monkeypatch: pytest.MonkeyPatch) -> No
|
||||
assert last_sort == [("seq", -1)]
|
||||
|
||||
|
||||
def test_save_state_assigns_sequence_and_inserts_in_one_transaction(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
created = _patch_client(monkeypatch)
|
||||
persistence = MongoDbFlowPersistence(CONN)
|
||||
|
||||
persistence.save_state("flow-1", "step", {"counter": 1})
|
||||
|
||||
assert created["client"].transactions_started == 1
|
||||
|
||||
|
||||
def test_basemodel_state_serialized_as_json(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
@@ -198,12 +231,13 @@ def test_dict_state_serializes_non_json_values(
|
||||
},
|
||||
)
|
||||
|
||||
assert persistence.load_state("flow-1") == {
|
||||
"when": "2026-01-02T03:04:05+00:00",
|
||||
"tags": ["a", "b"],
|
||||
"items": [1, 2],
|
||||
"nested": {"when": "2026-01-02T03:04:05Z"},
|
||||
}
|
||||
loaded = persistence.load_state("flow-1")
|
||||
|
||||
assert loaded is not None
|
||||
assert loaded["when"] == "2026-01-02T03:04:05+00:00"
|
||||
assert set(loaded["tags"]) == {"a", "b"}
|
||||
assert loaded["items"] == [1, 2]
|
||||
assert loaded["nested"] == {"when": "2026-01-02T03:04:05Z"}
|
||||
|
||||
|
||||
def test_missing_connection_string_raises(
|
||||
|
||||
Reference in New Issue
Block a user