Przeglądaj źródła

修复数据库连接问题

xueyiming 4 dni temu
rodzic
commit
94c5fc74d6

+ 5 - 1
supply_infra/db/repositories/demand_grade_plan_repo.py

@@ -191,7 +191,11 @@ class DemandGradePlanRepository(BaseRepository[DemandGradePlan]):
         graded = graded_names or set()
         from supply_infra.scheduler.plan_group_batch import resolve_demands_for_category_ids
 
-        demands = resolve_demands_for_category_ids(biz_dt, category_ids)
+        demands = resolve_demands_for_category_ids(
+            biz_dt,
+            category_ids,
+            session=self.session,
+        )
         created = 0
         for sort_order, demand in enumerate(demands, start=1):
             demand_name = str(demand["demand_name"])

+ 30 - 19
supply_infra/scheduler/plan_group_batch.py

@@ -3,6 +3,8 @@ from __future__ import annotations
 
 from typing import Any
 
+from sqlalchemy.orm import Session
+
 from supply_infra.db.repositories.demand_belong_category_repo import DemandBelongCategoryRepository
 from supply_infra.db.repositories.demand_belong_pool_rel_repo import DemandBelongPoolRelRepository
 from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
@@ -102,28 +104,37 @@ def _priority_sort_key(name: str, priority_index: dict) -> tuple:
 def resolve_demands_for_category_ids(
     biz_dt: str,
     category_ids: list[int],
+    *,
+    session: Session | None = None,
 ) -> list[dict[str, Any]]:
     """按分类节点解析全部需求池记录(pool_id + demand_name)。"""
     selected_ids = list(dict.fromkeys(int(value) for value in category_ids))
-    with get_session() as session:
-        belongs = DemandBelongCategoryRepository(session).list_by_category_ids(selected_ids)
-        pool_ids_by_belong = DemandBelongPoolRelRepository(session).get_pool_ids_by_belong_ids(
-            [int(row.id) for row in belongs]
-        )
-        pool_ids = sorted({pool_id for values in pool_ids_by_belong.values() for pool_id in values})
-        pool_repo = MultiDemandPoolDiRepository(session)
-        pool_rows = pool_repo.get_by_ids(pool_ids)
-        from agents.demand_grade_agent.tools.demand_priority import build_demand_priority_index
-
-        priority_index = build_demand_priority_index(pool_repo.list_by_biz_dt(biz_dt))
-        candidates: list[dict[str, Any]] = []
-        for row in pool_rows:
-            if row.biz_dt != biz_dt or not row.demand_name:
-                continue
-            candidates.append({
-                "pool_id": int(row.id),
-                "demand_name": str(row.demand_name),
-            })
+    if session is None:
+        with get_session() as owned_session:
+            return resolve_demands_for_category_ids(
+                biz_dt,
+                selected_ids,
+                session=owned_session,
+            )
+
+    belongs = DemandBelongCategoryRepository(session).list_by_category_ids(selected_ids)
+    pool_ids_by_belong = DemandBelongPoolRelRepository(session).get_pool_ids_by_belong_ids(
+        [int(row.id) for row in belongs]
+    )
+    pool_ids = sorted({pool_id for values in pool_ids_by_belong.values() for pool_id in values})
+    pool_repo = MultiDemandPoolDiRepository(session)
+    pool_rows = pool_repo.get_by_ids(pool_ids)
+    from agents.demand_grade_agent.tools.demand_priority import build_demand_priority_index
+
+    priority_index = build_demand_priority_index(pool_repo.list_by_biz_dt(biz_dt))
+    candidates: list[dict[str, Any]] = []
+    for row in pool_rows:
+        if row.biz_dt != biz_dt or not row.demand_name:
+            continue
+        candidates.append({
+            "pool_id": int(row.id),
+            "demand_name": str(row.demand_name),
+        })
 
     candidates.sort(
         key=lambda item: (

+ 32 - 0
tests/supply_infra/scheduler/test_plan_group_connection_reuse.py

@@ -0,0 +1,32 @@
+from __future__ import annotations
+
+from unittest.mock import Mock, patch
+
+from sqlalchemy.orm import Session
+
+from supply_infra.db.repositories.demand_grade_plan_repo import (
+    DemandGradePlanRepository,
+)
+
+
+def test_materialize_group_items_reuses_repository_session(monkeypatch) -> None:
+    session = Mock(spec=Session)
+    repository = DemandGradePlanRepository(session)
+    monkeypatch.setattr(repository, "group_item_count", lambda _group_id: 0)
+
+    with patch(
+        "supply_infra.scheduler.plan_group_batch.resolve_demands_for_category_ids",
+        return_value=[],
+    ) as resolve:
+        created = repository.materialize_group_items(
+            17,
+            "20260727",
+            [316, 317],
+        )
+
+    assert created == 0
+    resolve.assert_called_once_with(
+        "20260727",
+        [316, 317],
+        session=session,
+    )