o2o-negosium-original/agent/services/negotiation_service.py

119 lines
5.9 KiB
Python

"""NegotiationService — 협상 한 라운드 (실제 UCB Q-Table 학습 정책, H1).
흐름: 관측치 → build_state → 정책 로드(learning 스키마) → UCB 선택(propensity) → reward
→ Q-learning 온라인 갱신 + touched 셀 write-through → experience_logs 기록.
반복 호출하면 visit/Q 가 DB 에 누적되어 학습이 진행된다(같은 state 를 칠수록 탐색 보너스↓, Q 수렴).
대화형 /chat·step 체계·시퀀스 보상링크는 P5/P7. 여기는 단일 라운드 단위.
"""
import uuid
from common.enums import DBType, ErrorType
from common.database.db_session_manager import DB_SESSION_MNG
from common.logger import LOG
from negotiation.policies.base import EpisodeState, PolicyContext, Transition
from negotiation.policy.model_store import QTablePolicyStore
from negotiation.qtable.infra.repository.learning_repository import LearningRepository
from negotiation.qtable.domain.model.snapshot import NegotiationOutcome, NegotiationSnapshot
from negotiation.qtable.domain.service.reward_calculator import RewardCalculator
from negotiation.qtable.domain.service.state_calculator import build_state, state_index
from router.v1.negotiation.protocol import Req_NegotiationStep, Res_NegotiationStep, RewardView, StateView
from tenancy.registry import TenantEngine
class NegotiationService:
async def step(self, engine: TenantEngine, req: Req_NegotiationStep) -> Res_NegotiationStep:
res = Res_NegotiationStep(tenant_id=engine.tenant_id, company_id=engine.company_id)
# 1) 관측치 → snapshot
try:
outcome = NegotiationOutcome(req.outcome)
except ValueError:
res.result.SetResult(ErrorType.INVALID_REQUEST_DATA)
res.msg = f"outcome must be ongoing|success|failure, got {req.outcome!r}"
return res
# 앵커링값은 갑(KT/iMK)이 직접 입력한 값을 사용.
snap = NegotiationSnapshot(
revenue_amount=req.revenue_amount, distribution_code=req.distribution_code,
partner_count=req.partner_count, acceptance_ratio=req.acceptance_ratio,
input_price=req.input_price, anchor_price=req.anchor_price, target_price=req.target_price,
round_number=req.round_number, outcome=outcome,
)
# 2) 상태 산출 (config 주입)
try:
st = build_state(snap, engine.config.state)
idx = state_index(snap, engine.config.state)
except ValueError as ex:
res.result.SetResult(ErrorType.NEGO_INVALID_STEP)
res.msg = str(ex)
return res
# 3) 정책 로드 (learning 스키마 활성 버전) → UCB 선택
policy, version_id, repo = await QTablePolicyStore.load(engine)
episode = EpisodeState(used_action_ids=set(req.used_action_ids or []))
ctx = PolicyContext(state_index=idx, snapshot=snap, action_space_size=engine.action_space_size, episode=episode)
decision = policy.select(ctx)
decision.card_id = engine.mapper.get_card_id(decision.action_id)
# 4) 보상
reward = RewardCalculator(engine.config.reward).calculate(snap)
# 5) 학습: Q-learning 온라인 갱신 + touched 셀 영속화
updated_q = decision.q_value
if req.learn:
done = outcome != NegotiationOutcome.ONGOING
policy.update(Transition(state_index=idx, action_id=decision.action_id, reward=reward.total, done=done))
updated_q = float(policy.qtable.q[idx, decision.action_id])
try:
await QTablePolicyStore.persist_cell(repo, version_id, policy, idx, decision.action_id)
res.learned = True
except Exception as ex:
LOG.e_no_callstack(f"[NegotiationService] persist failed: {ex}")
# 6) 응답 채우기
session_id = req.session_id or str(uuid.uuid4())
res.session_id = session_id
res.state_index = idx
res.state = StateView(
revenue_idx=st.revenue_idx, distribution_idx=st.distribution_idx, partner_idx=st.partner_idx,
acceptance_idx=st.acceptance_idx, price_zone_idx=st.price_zone_idx,
)
res.action_id = decision.action_id
res.card_id = decision.card_id
res.propensity = decision.propensity
res.available_actions = decision.available_actions
res.reward = RewardView(
price_reward=reward.price_reward, end_reward=reward.end_reward,
penalty=reward.penalty, weight=reward.weight, total=reward.total,
)
res.policy = policy.name
res.q_value = decision.q_value
res.ucb_score = decision.ucb_score
res.updated_q = updated_q
res.visit_count = int(policy.qtable.visits[idx, decision.action_id])
# 7) experience_logs 기록
if req.log:
res.logged = await self._log(engine, session_id, idx, decision, snap, reward)
return res
async def _log(self, engine, session_id, idx, decision, snap, reward) -> bool:
repo = LearningRepository(engine.company_id)
data = {
"session_id": session_id, "state_index": idx, "action_id": decision.action_id,
"card_id": decision.card_id, "snapshot": snap.to_dict(), "propensity": decision.propensity,
"turn": snap.round_number, "available_actions": decision.available_actions,
"reward": reward.total, "done": snap.outcome != NegotiationOutcome.ONGOING,
"q_value_at_selection": decision.q_value, "ucb_score_at_selection": decision.ucb_score,
"settled_price": int(snap.input_price) if snap.outcome == NegotiationOutcome.SUCCESS else None,
}
try:
err = await DB_SESSION_MNG.execute_lambda_run([DBType.MAIN.value], [lambda s: repo.log_transition(s, data)])
return err == ErrorType.SUCCESS
except Exception as ex:
LOG.e_no_callstack(f"[NegotiationService] log failed: {ex}")
return False