omshrivastava commited on
Commit
10faff6
·
1 Parent(s): 107480e

backend: reserve restore slots and extract idle cleanup constants

Browse files
backend/session_manager.py CHANGED
@@ -124,6 +124,8 @@ class SessionCapacityError(Exception):
124
  MAX_SESSIONS: int = 200
125
  MAX_SESSIONS_PER_USER: int = 10
126
  DEFAULT_YOLO_COST_CAP_USD: float = 5.0
 
 
127
  SANDBOX_SHUTDOWN_CLEANUP_CONCURRENCY: int = 10
128
  SANDBOX_SHUTDOWN_CLEANUP_TIMEOUT_S: float = 60.0
129
 
@@ -208,6 +210,40 @@ class SessionManager:
208
  logger.info("Session initialized in %.2fs", t1 - t0)
209
  return tool_router, session
210
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
211
  def _serialize_messages(self, session: Session) -> list[dict[str, Any]]:
212
  return [msg.model_dump(mode="json") for msg in session.context_manager.items]
213
 
@@ -524,7 +560,7 @@ class SessionManager:
524
  """
525
  try:
526
  while True:
527
- await asyncio.sleep(600) # 10 minutes
528
  now = datetime.utcnow()
529
  to_unload: list[tuple[str, AgentSession]] = []
530
  async with self._lock:
@@ -534,12 +570,16 @@ class SessionManager:
534
  continue
535
  if getattr(agent_session, "is_processing", False):
536
  continue
537
- last = getattr(agent_session, "last_access", agent_session.created_at)
538
- if now - last > timedelta(hours=24):
539
- # Mark inactive and remove from registry under lock
 
 
 
 
 
540
  agent_session.is_active = False
541
  to_unload.append((sid, agent_session))
542
- del self.sessions[sid]
543
  except Exception as e:
544
  logger.debug("Skipping unload check for %s: %s", sid, e)
545
 
@@ -556,7 +596,24 @@ class SessionManager:
556
  sid,
557
  e,
558
  )
559
- # Cancel the running task if present
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
560
  try:
561
  if agent_session.task:
562
  agent_session.task.cancel()
@@ -631,146 +688,171 @@ class SessionManager:
631
  preload_sandbox: bool = True,
632
  ) -> AgentSession | None:
633
  """Return a live runtime session, lazily restoring it from Mongo."""
634
- async with self._lock:
635
- existing = self.sessions.get(session_id)
636
- if existing:
637
- if self._can_access_session(existing, user_id):
638
- self._update_hf_identity(
639
- existing,
640
- hf_token=hf_token,
641
- hf_username=hf_username,
642
- )
643
- self._restart_cpu_preload_if_token_recovered(
644
- existing,
645
- preload_sandbox=preload_sandbox,
646
- )
647
- existing.last_access = datetime.utcnow()
648
- return existing
649
- return None
650
 
651
- # Check capacity before restoring from persistence
652
- async with self._lock:
653
- active_count = self.active_session_count
654
- if active_count >= MAX_SESSIONS:
655
- logger.warning(
656
- "Cannot restore session %s: server at capacity (%d/%d)",
657
- session_id,
658
- active_count,
659
- MAX_SESSIONS,
660
- )
661
- return None
 
 
 
 
 
 
 
 
662
 
663
- store = self._store()
664
- loaded = await store.load_session(session_id)
665
- if not loaded:
666
- return None
 
 
 
 
 
667
 
668
- async with self._lock:
669
- existing = self.sessions.get(session_id)
670
- if existing:
671
- if self._can_access_session(existing, user_id):
672
- self._update_hf_identity(
673
- existing,
674
- hf_token=hf_token,
675
  hf_username=hf_username,
 
 
676
  )
677
- self._restart_cpu_preload_if_token_recovered(
678
- existing,
679
- preload_sandbox=preload_sandbox,
680
- )
681
- existing.last_access = datetime.utcnow()
682
- return existing
683
- return None
684
 
685
- meta = loaded.get("metadata") or {}
686
- owner = str(meta.get("user_id") or "")
687
- if user_id != "dev" and owner != "dev" and owner != user_id:
688
- return None
689
 
690
- await self._cleanup_persisted_sandbox(
691
- session_id,
692
- meta,
693
- hf_token=hf_token,
694
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
695
 
696
- from litellm import Message
 
 
 
 
697
 
698
- model = meta.get("model") or self.config.model_name
699
- event_queue: asyncio.Queue = asyncio.Queue()
700
- submission_queue: asyncio.Queue = asyncio.Queue()
701
- tool_router, session = await asyncio.to_thread(
702
- self._create_session_sync,
703
- session_id=session_id,
704
- user_id=owner or user_id,
705
- hf_username=hf_username,
706
- hf_token=hf_token,
707
- model=model,
708
- event_queue=event_queue,
709
- notification_destinations=meta.get("notification_destinations") or [],
710
- )
711
 
712
- restored_messages: list[Message] = []
713
- for raw in loaded.get("messages") or []:
714
- if not isinstance(raw, dict) or raw.get("role") == "system":
715
- continue
716
- try:
717
- restored_messages.append(Message.model_validate(raw))
718
- except Exception as e:
719
- logger.warning("Dropping malformed restored message: %s", e)
720
- if restored_messages:
721
- # Keep the freshly-rendered system prompt, then attach the durable
722
- # non-system context so tools/date/user context stay current.
723
- session.context_manager.items = [
724
- session.context_manager.items[0],
725
- *restored_messages,
726
- ]
727
-
728
- self._restore_pending_approval(session, meta.get("pending_approval") or [])
729
- session.turn_count = int(meta.get("turn_count") or 0)
730
- session.auto_approval_enabled = bool(meta.get("auto_approval_enabled", False))
731
- raw_cap = meta.get("auto_approval_cost_cap_usd")
732
- session.auto_approval_cost_cap_usd = (
733
- float(raw_cap) if isinstance(raw_cap, int | float) else None
734
- )
735
- session.auto_approval_estimated_spend_usd = float(
736
- meta.get("auto_approval_estimated_spend_usd") or 0.0
737
- )
738
 
739
- created_at = meta.get("created_at")
740
- if not isinstance(created_at, datetime):
741
- created_at = datetime.utcnow()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
742
 
743
- agent_session = AgentSession(
744
- session_id=session_id,
745
- session=session,
746
- tool_router=tool_router,
747
- submission_queue=submission_queue,
748
- user_id=owner or user_id,
749
- hf_username=hf_username,
750
- hf_token=hf_token,
751
- created_at=created_at,
752
- is_active=True,
753
- is_processing=False,
754
- claude_counted=bool(meta.get("claude_counted")),
755
- title=meta.get("title"),
756
- )
757
- started = await self._start_agent_session(
758
- agent_session=agent_session,
759
- event_queue=event_queue,
760
- tool_router=tool_router,
761
- )
762
- if started is not agent_session:
763
- self._update_hf_identity(
764
- started,
765
- hf_token=hf_token,
766
  hf_username=hf_username,
 
 
 
 
 
 
767
  )
768
- started.last_access = datetime.utcnow()
769
- return started
770
- if preload_sandbox:
771
- self._start_cpu_sandbox_preload(agent_session)
772
- logger.info("Restored session %s for user %s", session_id, owner or user_id)
773
- return agent_session
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
774
 
775
  async def create_session(
776
  self,
@@ -838,41 +920,64 @@ class SessionManager:
838
  self.sessions[session_id] = placeholder
839
 
840
  # Run blocking constructors in a thread to keep the event loop responsive.
841
- tool_router, session = await asyncio.to_thread(
842
- self._create_session_sync,
843
- session_id=session_id,
844
- user_id=user_id,
845
- hf_username=hf_username,
846
- hf_token=hf_token,
847
- model=model,
848
- event_queue=event_queue,
849
- )
 
 
850
 
851
- # Create wrapper with the real session resources and replace the
852
- # placeholder in _start_agent_session.
853
- agent_session = AgentSession(
854
- session_id=session_id,
855
- session=session,
856
- tool_router=tool_router,
857
- submission_queue=submission_queue,
858
- user_id=user_id,
859
- hf_username=hf_username,
860
- hf_token=hf_token,
861
- )
862
 
863
- await self._start_agent_session(
864
- agent_session=agent_session,
865
- event_queue=event_queue,
866
- tool_router=tool_router,
867
- )
868
- await self.persist_session_snapshot(agent_session, runtime_state="idle")
869
- self._start_cpu_sandbox_preload(agent_session)
870
 
871
- if is_pro is not None and user_id and user_id != "dev":
872
- await self._track_pro_status(agent_session, is_pro=is_pro)
873
 
874
- logger.info(f"Created session {session_id} for user {user_id}")
875
- return session_id
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
876
 
877
  async def _track_pro_status(
878
  self, agent_session: AgentSession, *, is_pro: bool
 
124
  MAX_SESSIONS: int = 200
125
  MAX_SESSIONS_PER_USER: int = 10
126
  DEFAULT_YOLO_COST_CAP_USD: float = 5.0
127
+ INACTIVE_SESSION_SWEEP_INTERVAL_SECONDS: int = 600
128
+ INACTIVE_SESSION_IDLE_THRESHOLD: timedelta = timedelta(hours=24)
129
  SANDBOX_SHUTDOWN_CLEANUP_CONCURRENCY: int = 10
130
  SANDBOX_SHUTDOWN_CLEANUP_TIMEOUT_S: float = 60.0
131
 
 
210
  logger.info("Session initialized in %.2fs", t1 - t0)
211
  return tool_router, session
212
 
213
+ def _make_reserved_session(
214
+ self,
215
+ *,
216
+ session_id: str,
217
+ user_id: str,
218
+ hf_username: str | None,
219
+ hf_token: str | None,
220
+ submission_queue: asyncio.Queue,
221
+ ) -> AgentSession:
222
+ """Create a placeholder session that reserves capacity under lock."""
223
+ return AgentSession(
224
+ session_id=session_id,
225
+ session=None,
226
+ tool_router=None,
227
+ submission_queue=submission_queue,
228
+ user_id=user_id,
229
+ hf_username=hf_username,
230
+ hf_token=hf_token,
231
+ is_active=True,
232
+ )
233
+
234
+ async def _release_reserved_session_slot(
235
+ self,
236
+ session_id: str,
237
+ reserved_session: AgentSession | None = None,
238
+ ) -> None:
239
+ """Remove a reserved placeholder if it is still present."""
240
+ async with self._lock:
241
+ current = self.sessions.get(session_id)
242
+ if current is None:
243
+ return
244
+ if current is reserved_session or getattr(current, "session", None) is None:
245
+ self.sessions.pop(session_id, None)
246
+
247
  def _serialize_messages(self, session: Session) -> list[dict[str, Any]]:
248
  return [msg.model_dump(mode="json") for msg in session.context_manager.items]
249
 
 
560
  """
561
  try:
562
  while True:
563
+ await asyncio.sleep(INACTIVE_SESSION_SWEEP_INTERVAL_SECONDS)
564
  now = datetime.utcnow()
565
  to_unload: list[tuple[str, AgentSession]] = []
566
  async with self._lock:
 
570
  continue
571
  if getattr(agent_session, "is_processing", False):
572
  continue
573
+ last = getattr(
574
+ agent_session,
575
+ "last_access",
576
+ agent_session.created_at,
577
+ )
578
+ if now - last > INACTIVE_SESSION_IDLE_THRESHOLD:
579
+ # Mark inactive, but keep it resident until the
580
+ # snapshot has been persisted successfully.
581
  agent_session.is_active = False
582
  to_unload.append((sid, agent_session))
 
583
  except Exception as e:
584
  logger.debug("Skipping unload check for %s: %s", sid, e)
585
 
 
596
  sid,
597
  e,
598
  )
599
+ # Keep the session in memory so the next cleanup cycle
600
+ # can retry persistence and so callers can still inspect
601
+ # the session state.
602
+ agent_session.is_active = True
603
+ continue
604
+
605
+ removed = False
606
+ async with self._lock:
607
+ current = self.sessions.get(sid)
608
+ if current is agent_session:
609
+ self.sessions.pop(sid, None)
610
+ removed = True
611
+
612
+ if not removed:
613
+ # Session was replaced or revived while we were persisting.
614
+ continue
615
+
616
+ # Cancel the running task if present once the snapshot is safe.
617
  try:
618
  if agent_session.task:
619
  agent_session.task.cancel()
 
688
  preload_sandbox: bool = True,
689
  ) -> AgentSession | None:
690
  """Return a live runtime session, lazily restoring it from Mongo."""
691
+ submission_queue: asyncio.Queue = asyncio.Queue()
692
+ event_queue: asyncio.Queue = asyncio.Queue()
693
+ reserved_session: AgentSession | None = None
694
+ should_release_reserved_slot = False
 
 
 
 
 
 
 
 
 
 
 
 
695
 
696
+ try:
697
+ async with self._lock:
698
+ existing = self.sessions.get(session_id)
699
+ if existing:
700
+ if getattr(existing, "session", None) is None:
701
+ return None
702
+ if self._can_access_session(existing, user_id):
703
+ self._update_hf_identity(
704
+ existing,
705
+ hf_token=hf_token,
706
+ hf_username=hf_username,
707
+ )
708
+ self._restart_cpu_preload_if_token_recovered(
709
+ existing,
710
+ preload_sandbox=preload_sandbox,
711
+ )
712
+ existing.last_access = datetime.utcnow()
713
+ return existing
714
+ return None
715
 
716
+ active_count = self.active_session_count
717
+ if active_count >= MAX_SESSIONS:
718
+ logger.warning(
719
+ "Cannot restore session %s: server at capacity (%d/%d)",
720
+ session_id,
721
+ active_count,
722
+ MAX_SESSIONS,
723
+ )
724
+ return None
725
 
726
+ reserved_session = self._make_reserved_session(
727
+ session_id=session_id,
728
+ user_id=user_id,
 
 
 
 
729
  hf_username=hf_username,
730
+ hf_token=hf_token,
731
+ submission_queue=submission_queue,
732
  )
733
+ self.sessions[session_id] = reserved_session
734
+ should_release_reserved_slot = True
 
 
 
 
 
735
 
736
+ store = self._store()
737
+ loaded = await store.load_session(session_id)
738
+ if not loaded:
739
+ return None
740
 
741
+ async with self._lock:
742
+ existing = self.sessions.get(session_id)
743
+ if existing is not reserved_session:
744
+ if existing and getattr(existing, "session", None) is not None:
745
+ if self._can_access_session(existing, user_id):
746
+ self._update_hf_identity(
747
+ existing,
748
+ hf_token=hf_token,
749
+ hf_username=hf_username,
750
+ )
751
+ self._restart_cpu_preload_if_token_recovered(
752
+ existing,
753
+ preload_sandbox=preload_sandbox,
754
+ )
755
+ existing.last_access = datetime.utcnow()
756
+ should_release_reserved_slot = False
757
+ return existing
758
+ return None
759
+
760
+ meta = loaded.get("metadata") or {}
761
+ owner = str(meta.get("user_id") or "")
762
+ if user_id != "dev" and owner != "dev" and owner != user_id:
763
+ return None
764
 
765
+ await self._cleanup_persisted_sandbox(
766
+ session_id,
767
+ meta,
768
+ hf_token=hf_token,
769
+ )
770
 
771
+ from litellm import Message
 
 
 
 
 
 
 
 
 
 
 
 
772
 
773
+ model = meta.get("model") or self.config.model_name
774
+ tool_router, session = await asyncio.to_thread(
775
+ self._create_session_sync,
776
+ session_id=session_id,
777
+ user_id=owner or user_id,
778
+ hf_username=hf_username,
779
+ hf_token=hf_token,
780
+ model=model,
781
+ event_queue=event_queue,
782
+ notification_destinations=meta.get("notification_destinations") or [],
783
+ )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
784
 
785
+ restored_messages: list[Message] = []
786
+ for raw in loaded.get("messages") or []:
787
+ if not isinstance(raw, dict) or raw.get("role") == "system":
788
+ continue
789
+ try:
790
+ restored_messages.append(Message.model_validate(raw))
791
+ except Exception as e:
792
+ logger.warning("Dropping malformed restored message: %s", e)
793
+ if restored_messages:
794
+ # Keep the freshly-rendered system prompt, then attach the durable
795
+ # non-system context so tools/date/user context stay current.
796
+ session.context_manager.items = [
797
+ session.context_manager.items[0],
798
+ *restored_messages,
799
+ ]
800
+
801
+ self._restore_pending_approval(session, meta.get("pending_approval") or [])
802
+ session.turn_count = int(meta.get("turn_count") or 0)
803
+ session.auto_approval_enabled = bool(
804
+ meta.get("auto_approval_enabled", False)
805
+ )
806
+ raw_cap = meta.get("auto_approval_cost_cap_usd")
807
+ session.auto_approval_cost_cap_usd = (
808
+ float(raw_cap) if isinstance(raw_cap, int | float) else None
809
+ )
810
+ session.auto_approval_estimated_spend_usd = float(
811
+ meta.get("auto_approval_estimated_spend_usd") or 0.0
812
+ )
813
 
814
+ created_at = meta.get("created_at")
815
+ if not isinstance(created_at, datetime):
816
+ created_at = datetime.utcnow()
817
+
818
+ agent_session = AgentSession(
819
+ session_id=session_id,
820
+ session=session,
821
+ tool_router=tool_router,
822
+ submission_queue=submission_queue,
823
+ user_id=owner or user_id,
 
 
 
 
 
 
 
 
 
 
 
 
 
824
  hf_username=hf_username,
825
+ hf_token=hf_token,
826
+ created_at=created_at,
827
+ is_active=True,
828
+ is_processing=False,
829
+ claude_counted=bool(meta.get("claude_counted")),
830
+ title=meta.get("title"),
831
  )
832
+ started = await self._start_agent_session(
833
+ agent_session=agent_session,
834
+ event_queue=event_queue,
835
+ tool_router=tool_router,
836
+ )
837
+ if started is not agent_session:
838
+ self._update_hf_identity(
839
+ started,
840
+ hf_token=hf_token,
841
+ hf_username=hf_username,
842
+ )
843
+ started.last_access = datetime.utcnow()
844
+ should_release_reserved_slot = False
845
+ return started
846
+ if preload_sandbox:
847
+ self._start_cpu_sandbox_preload(agent_session)
848
+ logger.info("Restored session %s for user %s", session_id, owner or user_id)
849
+ should_release_reserved_slot = False
850
+ return agent_session
851
+ except Exception:
852
+ raise
853
+ finally:
854
+ if should_release_reserved_slot:
855
+ await self._release_reserved_session_slot(session_id, reserved_session)
856
 
857
  async def create_session(
858
  self,
 
920
  self.sessions[session_id] = placeholder
921
 
922
  # Run blocking constructors in a thread to keep the event loop responsive.
923
+ agent_session: AgentSession | None = None
924
+ try:
925
+ tool_router, session = await asyncio.to_thread(
926
+ self._create_session_sync,
927
+ session_id=session_id,
928
+ user_id=user_id,
929
+ hf_username=hf_username,
930
+ hf_token=hf_token,
931
+ model=model,
932
+ event_queue=event_queue,
933
+ )
934
 
935
+ # Create wrapper with the real session resources and replace the
936
+ # placeholder in _start_agent_session.
937
+ agent_session = AgentSession(
938
+ session_id=session_id,
939
+ session=session,
940
+ tool_router=tool_router,
941
+ submission_queue=submission_queue,
942
+ user_id=user_id,
943
+ hf_username=hf_username,
944
+ hf_token=hf_token,
945
+ )
946
 
947
+ await self._start_agent_session(
948
+ agent_session=agent_session,
949
+ event_queue=event_queue,
950
+ tool_router=tool_router,
951
+ )
952
+ await self.persist_session_snapshot(agent_session, runtime_state="idle")
953
+ self._start_cpu_sandbox_preload(agent_session)
954
 
955
+ if is_pro is not None and user_id and user_id != "dev":
956
+ await self._track_pro_status(agent_session, is_pro=is_pro)
957
 
958
+ logger.info(f"Created session {session_id} for user {user_id}")
959
+ return session_id
960
+ except Exception:
961
+ cleanup_task: asyncio.Task | None = None
962
+ async with self._lock:
963
+ current = self.sessions.get(session_id)
964
+ if current and (
965
+ current is agent_session
966
+ or getattr(current, "session", None) is None
967
+ ):
968
+ self.sessions.pop(session_id, None)
969
+ if agent_session is not None:
970
+ agent_session.is_active = False
971
+ cleanup_task = agent_session.task
972
+ if cleanup_task and not cleanup_task.done():
973
+ cleanup_task.cancel()
974
+ try:
975
+ await cleanup_task
976
+ except asyncio.CancelledError:
977
+ pass
978
+ except Exception:
979
+ pass
980
+ raise
981
 
982
  async def _track_pro_status(
983
  self, agent_session: AgentSession, *, is_pro: bool
tests/unit/test_session_capacity.py CHANGED
@@ -1,4 +1,5 @@
1
  import asyncio
 
2
  import sys
3
  from pathlib import Path
4
  from types import SimpleNamespace
@@ -6,31 +7,57 @@ from types import SimpleNamespace
6
  import pytest
7
  import types
8
 
9
- # Prevent importing heavy third-party modules when importing the backend module.
10
- for _mod in ("litellm", "fastmcp", "thefuzz", "huggingface_hub"):
11
- if _mod not in sys.modules:
12
- m = types.ModuleType(_mod)
13
- # fastmcp is imported as `from fastmcp import Client` in some codepaths
14
- if _mod == "fastmcp":
15
- class _DummyClient:
16
- pass
17
 
18
- setattr(m, "Client", _DummyClient)
19
- sys.modules[_mod] = m
20
 
21
- _BACKEND_DIR = Path(__file__).resolve().parent.parent.parent / "backend"
22
- if str(_BACKEND_DIR) not in sys.path:
23
- sys.path.insert(0, str(_BACKEND_DIR))
 
 
 
 
 
 
 
 
 
24
 
25
- from session_manager import SessionManager, AgentSession, MAX_SESSIONS
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
26
 
27
 
28
  @pytest.mark.asyncio
29
- async def test_restore_denied_when_at_capacity(caplog):
30
- manager = SessionManager()
31
  # Fill in-memory sessions up to MAX_SESSIONS
32
- for i in range(MAX_SESSIONS):
33
- manager.sessions[str(i)] = AgentSession(
34
  session_id=str(i),
35
  session=object(),
36
  tool_router=None,
@@ -53,8 +80,8 @@ async def test_restore_denied_when_at_capacity(caplog):
53
 
54
 
55
  @pytest.mark.asyncio
56
- async def test_restore_allowed_under_capacity(monkeypatch):
57
- manager = SessionManager()
58
  manager.sessions.clear()
59
 
60
  class DummyStore:
@@ -87,4 +114,39 @@ async def test_restore_allowed_under_capacity(monkeypatch):
87
 
88
  res = await manager.ensure_session_loaded("restored", user_id="test")
89
  assert res is not None
90
- assert res.session is not None
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  import asyncio
2
+ import importlib
3
  import sys
4
  from pathlib import Path
5
  from types import SimpleNamespace
 
7
  import pytest
8
  import types
9
 
10
+ _BACKEND_DIR = Path(__file__).resolve().parent.parent.parent / "backend"
 
 
 
 
 
 
 
11
 
 
 
12
 
13
+ @pytest.fixture
14
+ def session_manager_module(monkeypatch):
15
+ """Import backend.session_manager with temporary dependency stubs.
16
+
17
+ The stubs are inserted with monkeypatch so they are restored after the test,
18
+ and the imported session_manager module is removed from sys.modules to avoid
19
+ leaking the stubbed import state into other tests.
20
+ """
21
+ with monkeypatch.context() as m:
22
+ m.syspath_prepend(str(_BACKEND_DIR))
23
+
24
+ litellm_stub = types.ModuleType("litellm")
25
 
26
+ class _DummyMessage:
27
+ @staticmethod
28
+ def model_validate(raw):
29
+ return SimpleNamespace(**raw)
30
+
31
+ setattr(litellm_stub, "Message", _DummyMessage)
32
+ m.setitem(sys.modules, "litellm", litellm_stub)
33
+
34
+ fastmcp_stub = types.ModuleType("fastmcp")
35
+
36
+ class _DummyClient:
37
+ pass
38
+
39
+ setattr(fastmcp_stub, "Client", _DummyClient)
40
+ m.setitem(sys.modules, "fastmcp", fastmcp_stub)
41
+
42
+ m.setitem(sys.modules, "thefuzz", types.ModuleType("thefuzz"))
43
+ m.setitem(
44
+ sys.modules,
45
+ "huggingface_hub",
46
+ types.ModuleType("huggingface_hub"),
47
+ )
48
+
49
+ module = importlib.import_module("session_manager")
50
+ yield module
51
+
52
+ sys.modules.pop("session_manager", None)
53
 
54
 
55
  @pytest.mark.asyncio
56
+ async def test_restore_denied_when_at_capacity(session_manager_module, caplog):
57
+ manager = session_manager_module.SessionManager()
58
  # Fill in-memory sessions up to MAX_SESSIONS
59
+ for i in range(session_manager_module.MAX_SESSIONS):
60
+ manager.sessions[str(i)] = session_manager_module.AgentSession(
61
  session_id=str(i),
62
  session=object(),
63
  tool_router=None,
 
80
 
81
 
82
  @pytest.mark.asyncio
83
+ async def test_restore_allowed_under_capacity(session_manager_module, monkeypatch):
84
+ manager = session_manager_module.SessionManager()
85
  manager.sessions.clear()
86
 
87
  class DummyStore:
 
114
 
115
  res = await manager.ensure_session_loaded("restored", user_id="test")
116
  assert res is not None
117
+ assert res.session is not None
118
+
119
+
120
+ @pytest.mark.asyncio
121
+ async def test_restore_rolls_back_placeholder_on_load_failure(session_manager_module):
122
+ manager = session_manager_module.SessionManager()
123
+
124
+ class FailingStore:
125
+ enabled = True
126
+
127
+ async def load_session(self, sid):
128
+ raise RuntimeError("load failed")
129
+
130
+ manager.persistence_store = FailingStore()
131
+
132
+ with pytest.raises(RuntimeError, match="load failed"):
133
+ await manager.ensure_session_loaded("restored", user_id="test")
134
+
135
+ assert manager.sessions == {}
136
+
137
+
138
+ @pytest.mark.asyncio
139
+ async def test_create_session_rolls_back_placeholder_on_failure(
140
+ session_manager_module, monkeypatch
141
+ ):
142
+ manager = session_manager_module.SessionManager()
143
+
144
+ def fake_create_session_sync(**kwargs):
145
+ raise RuntimeError("boom")
146
+
147
+ monkeypatch.setattr(manager, "_create_session_sync", fake_create_session_sync)
148
+
149
+ with pytest.raises(RuntimeError, match="boom"):
150
+ await manager.create_session(user_id="u1")
151
+
152
+ assert manager.sessions == {}