"""learning/qtable 관리 API (Chat_server 14개 API 보존, tenant 스코프). 엔드포인트: /q-table/versions·switch·current, /experience-logs, /reset-learning, /reset-all, /invalidate-session, /train. 모두 X-Tenant-ID 헤더로 company_id 격리. """ from typing import Optional from fastapi import APIRouter, Depends from pydantic import BaseModel from common.enums import ErrorType from negotiation.policies.base import Transition from negotiation.policies.qtable_policy import UCBQTablePolicy from negotiation.policy.model_store import QTablePolicyStore from negotiation.qtable.domain.model.q_table import QTable from negotiation.qtable.infra.repository.learning_repository import LearningRepository from router.deps import get_tenant_engine from tenancy.registry import TenantEngine router = APIRouter(prefix="/v1", tags=["Learning / Q-Table"], responses={404: {"description": "Not found"}}) def _repo(engine: TenantEngine) -> LearningRepository: return LearningRepository(engine.company_id) @router.get("/q-table/versions", summary="Q-Table 버전 목록") async def list_versions(engine: TenantEngine = Depends(get_tenant_engine)): repo = _repo(engine) err, versions = await repo.read(lambda s: repo.list_versions(s)) return {"company_id": engine.company_id, "versions": [ {"version_name": v.version_name, "is_active": v.is_active, "scope": v.scope, "state_space_size": v.state_space_size, "action_space_size": v.action_space_size, "epochs": v.epochs, "created_at": str(v.created_at)} for v in (versions or [])]} class SwitchReq(BaseModel): version_name: str @router.post("/q-table/switch", summary="활성 Q-Table 버전 전환") async def switch_version(req: SwitchReq, engine: TenantEngine = Depends(get_tenant_engine)): err = await _repo(engine).set_active_version(req.version_name) ok = err == ErrorType.SUCCESS return {"success": ok, "active_version": req.version_name if ok else None, "desc": err.name} @router.get("/q-table/current", summary="현재 활성 Q-Table 상태") async def current_qtable(engine: TenantEngine = Depends(get_tenant_engine)): repo = _repo(engine) err, active = await repo.read(lambda s: repo.get_active_version(s)) if not active: return {"company_id": engine.company_id, "active_version": None, "q_value_rows": 0} rows = await repo.qvalue_count(active.version_id) return {"company_id": engine.company_id, "active_version": active.version_name, "state_space_size": active.state_space_size, "action_space_size": active.action_space_size, "q_value_rows": rows} @router.get("/experience-logs", summary="최근 경험 로그") async def experience_logs(limit: int = 50, engine: TenantEngine = Depends(get_tenant_engine)): repo = _repo(engine) err, total = await repo.read(lambda s: repo.count_experience(s)) logs = await repo.recent_experience(limit=min(limit, 500)) return {"company_id": engine.company_id, "total": total, "logs": logs} @router.post("/reset-learning", summary="[관리] 학습 데이터 초기화(자사 스코프)") async def reset_learning(engine: TenantEngine = Depends(get_tenant_engine)): err = await _repo(engine).reset_learning() return {"success": err == ErrorType.SUCCESS, "company_id": engine.company_id, "desc": err.name} @router.post("/reset-all", summary="[관리] 완전 초기화(자사 스코프, 타테넌트 무영향)") async def reset_all(engine: TenantEngine = Depends(get_tenant_engine)): err = await _repo(engine).reset_all() return {"success": err == ErrorType.SUCCESS, "company_id": engine.company_id, "desc": err.name} class InvalidateReq(BaseModel): session_id: str reason: Optional[str] = "invalidated" @router.post("/invalidate-session", summary="세션 학습 데이터 무효화") async def invalidate_session(req: InvalidateReq, engine: TenantEngine = Depends(get_tenant_engine)): repo = _repo(engine) err = await repo.write_one(lambda s: repo.invalidate_by_session(s, req.session_id, req.reason)) return {"success": err == ErrorType.SUCCESS, "session_id": req.session_id, "desc": err.name} class TrainReq(BaseModel): epochs: int = 1 @router.post("/train", summary="오프라인 Q-Learning 학습(경험로그 기반)") async def train(req: TrainReq, engine: TenantEngine = Depends(get_tenant_engine)): repo = _repo(engine) transitions = await repo.load_transitions() if not transitions: return {"success": True, "trained_transitions": 0, "note": "경험 로그 없음 (먼저 /chat 진행)"} S, A = engine.state_space_size, engine.action_space_size qt = QTable(S, A, learning_rate=engine.config.policy.learning_rate, discount_factor=engine.config.policy.gamma) # 활성 버전에서 현재 Q 적재 version_id = await repo.get_or_create_active_version( state_space_size=S, action_space_size=A, learning_rate=engine.config.policy.learning_rate, discount_factor=engine.config.policy.gamma) qrows, vrows = await repo.load_cells(version_id) for st, a, q in qrows: if 0 <= st < S and 0 <= a < A: qt.q[st, a] = q policy = UCBQTablePolicy(qt, mark_visits=False) updates = 0 touched = set() for _ in range(max(1, req.epochs)): for st, a, r, ns, done in transitions: if st is None or a is None or not (0 <= st < S and 0 <= a < A): continue policy.update(Transition(state_index=st, action_id=a, reward=r, next_state_index=ns, done=bool(done))) touched.add((st, a)) updates += 1 # 갱신된 셀 영속화 for (st, a) in touched: await QTablePolicyStore.persist_cell(repo, version_id, policy, st, a) return {"success": True, "trained_transitions": len(transitions), "epochs": req.epochs, "updates": updates, "cells_persisted": len(touched), "company_id": engine.company_id} @router.get("/verification-report", summary="Q-Table 검증 리포트(요약)") async def verification_report(engine: TenantEngine = Depends(get_tenant_engine)): repo = _repo(engine) err, active = await repo.read(lambda s: repo.get_active_version(s)) err, total_exp = await repo.read(lambda s: repo.count_experience(s)) if not active: return {"company_id": engine.company_id, "active_version": None, "experience_total": total_exp} qrows, vrows = await repo.load_cells(active.version_id) nonzero = [(s, a, q) for s, a, q in qrows if q != 0] top = sorted(nonzero, key=lambda x: -x[2])[:10] return { "company_id": engine.company_id, "active_version": active.version_name, "experience_total": total_exp, "q_value_rows": len(qrows), "nonzero_q": len(nonzero), "visited_cells": len(vrows), "top_q": [{"state_index": s, "action_id": a, "q_value": round(q, 4)} for s, a, q in top], }