from fastapi import Depends from common.database.db_session_manager import DB_SESSION_MNG from common.database.model.models import quotations from common.enums import DBWRType, ErrorType, SupplierType from crud.learning_crud import ILearningCRUD, LearningCRUD from router.v1.learning.protocol import ( AnchoringCell, AnchoringHistoryRow, CardPerformanceRow, LearningKpi, Res_AnchoringStatus, Res_LearningStatus, ) HISTORY_LIMIT = 50 # 앵커링 조정 이력 표시 개수 — 한 화면에서 훑는 용도 _SUPPLIER_TYPE_LABEL = { SupplierType.NONE.value: "미지정", SupplierType.DISTRIBUTION.value: "유통", SupplierType.MANUFACTURE.value: "제조", SupplierType.SOLE_AGENCY.value: "총판", } class LearningService: """협상 학습 현황 — 협상카드 학습(agent Q-learning)과 앵커링 조정 이력을 읽어 보여준다. 두 값 모두 negodata 가 만드는 값이 아니라 agent·anchoring 서비스가 쌓은 결과다(읽기 전용). """ def __init__(self, crud: ILearningCRUD = Depends(LearningCRUD)): self.crud = crud async def get_learning_status(self, company_id: str) -> Res_LearningStatus: res = Res_LearningStatus() # learning·anchoring 의 company_id 는 agent 가 테넌트 키를 그대로 넣는 문자열 컬럼이다(UUID 타입 아님). cid = str(company_id) summary = await self._read(lambda s: self.crud.learning_summary(s, cid), default=(0, 0, 0, None)) sessions, records, settled, last_at = summary res.kpi = LearningKpi( learned_sessions=int(sessions or 0), records=int(records or 0), settled_sessions=int(settled or 0), settle_rate=round((settled or 0) / sessions, 3) if sessions else 0.0, last_learned_at=last_at, ) rows = await self._read(lambda s: self.crud.card_performance(s, cid), default=[]) names = await self._read(lambda s: self.crud.card_names(s), default=[]) name_map = {str(number): (name, is_wild) for number, name, is_wild in names if number} for card_id, used_sessions, uses, avg_reward, settled_sessions in rows: number = str(card_id) name, is_wild = name_map.get(number, (None, 0)) res.cards.append(CardPerformanceRow( number=number, name=name, type="wild" if is_wild else "nego", used_sessions=int(used_sessions or 0), uses=int(uses or 0), avg_reward=round(float(avg_reward), 3) if avg_reward is not None else 0.0, settled_sessions=int(settled_sessions or 0), settle_rate=round((settled_sessions or 0) / used_sessions, 3) if used_sessions else 0.0, )) return res async def get_anchoring_status(self, company_id: str) -> Res_AnchoringStatus: res = Res_AnchoringStatus() cid = str(company_id) cells = await self._read(lambda s: self.crud.anchoring_current(s, cid), default=[]) for supplier_type, price_range_index, value, adjusted_at in cells: res.cells.append(AnchoringCell( supplier_type=int(supplier_type or 0), supplier_type_label=_SUPPLIER_TYPE_LABEL.get(int(supplier_type or 0), "미지정"), price_range_index=int(price_range_index or 0), anchoring_value=float(value) if value is not None else 0.0, last_adjusted_at=adjusted_at, )) history = await self._read(lambda s: self.crud.anchoring_history(s, cid, HISTORY_LIMIT), default=[]) for st, pri, before, after, sample, success, rate, created_at in history: res.history.append(AnchoringHistoryRow( supplier_type_label=_SUPPLIER_TYPE_LABEL.get(int(st or 0), "미지정"), price_range_index=int(pri or 0), value_before=float(before) if before is not None else 0.0, value_after=float(after) if after is not None else 0.0, sample_count=int(sample or 0), success_count=int(success or 0), success_rate=round(float(rate), 3) if rate is not None else 0.0, created_at=created_at, )) res.adjusted_count = len(res.history) return res async def _read(self, fn, default): """crud 한 건 실행 — 실패해도 화면은 떠야 하므로 기본값으로 떨어진다.""" err, rows = await DB_SESSION_MNG.execute_lambda(quotations.DBType(), DBWRType.DB_READ.value, fn) return rows if err == ErrorType.SUCCESS else default