o2o-negosium-original/agent/tests/test_p3_learning_schema.py

127 lines
5.4 KiB
Python

"""P3 검증 (계획서 P3 _검증_).
1. learning 스키마 테이블/제약 무결성 (create + 컬럼 propensity/turn/company_id 등 존재).
2. 2테넌트 동일 version_name 공존 (UNIQUE(company_id, version_name) 덕분).
3. company_id 자동 주입 + 위조 방지(log_transition 이 company_id 강제).
4. reset_all 이 자사 데이터만 삭제, 타테넌트 무영향 (파괴 테스트).
5. experience_logs 신규 컬럼(propensity/turn/available_actions/settled_price) 채워짐.
DB 미가용 시 db_engine 픽스처가 skip 한다.
"""
import uuid
import pytest
from common.database.db_session_manager import DB_SESSION_MNG
from common.enums import DBType, DBWRType, ErrorType
from negotiation.qtable.domain.model.snapshot import NegotiationOutcome, NegotiationSnapshot
from negotiation.qtable.infra.repository.learning_repository import LearningRepository
COMPANY_A = "company-aaaa"
COMPANY_B = "company-bbbb"
async def _read(func):
return await DB_SESSION_MNG.execute_lambda(DBType.MAIN.value, DBWRType.DB_READ.value, func)
@pytest.mark.asyncio
async def test_two_tenants_same_version_name_coexist(db_engine):
repo_a = LearningRepository(COMPANY_A)
repo_b = LearningRepository(COMPANY_B)
# 동일 version_name 을 두 회사가 각각 생성 — 공존해야 한다.
err_a = await DB_SESSION_MNG.execute_lambda_run(
[DBType.MAIN.value],
[lambda s: repo_a.create_version(s, version_name="v000", scope=2, state_space_size=162, action_space_size=9, is_active=True)],
)
err_b = await DB_SESSION_MNG.execute_lambda_run(
[DBType.MAIN.value],
[lambda s: repo_b.create_version(s, version_name="v000", scope=2, state_space_size=162, action_space_size=9, is_active=True)],
)
assert err_a == ErrorType.SUCCESS
assert err_b == ErrorType.SUCCESS
err, va = await _read(lambda s: repo_a.get_version_by_name(s, "v000"))
assert err == ErrorType.SUCCESS and va is not None and va.company_id == COMPANY_A
err, vb = await _read(lambda s: repo_b.get_version_by_name(s, "v000"))
assert vb is not None and vb.company_id == COMPANY_B
assert va.version_id != vb.version_id
@pytest.mark.asyncio
async def test_duplicate_version_same_company_rejected(db_engine):
repo = LearningRepository(COMPANY_A)
mk = lambda s: repo.create_version(s, version_name="dup", scope=2, state_space_size=162, action_space_size=9)
assert await DB_SESSION_MNG.execute_lambda_run([DBType.MAIN.value], [mk]) == ErrorType.SUCCESS
# 같은 회사 + 같은 version_name → 유니크 위반
err = await DB_SESSION_MNG.execute_lambda_run([DBType.MAIN.value], [mk])
assert err == ErrorType.DB_ALREADY_SAME_KEY
@pytest.mark.asyncio
async def test_log_transition_forces_company_id_and_new_columns(db_engine):
repo = LearningRepository(COMPANY_A)
snap = NegotiationSnapshot(
revenue_amount=5_000_000, distribution_code="A", partner_count=1, acceptance_ratio=0.05,
input_price=900, anchor_price=800, target_price=1000, round_number=2, outcome=NegotiationOutcome.ONGOING,
)
data = {
"company_id": "ATTACKER", # 위조 시도 — repo 가 자사 company_id 로 덮어써야 함
"state_index": 55, "action_id": 3, "card_id": "NGC-A004",
"snapshot": snap.to_dict(),
"propensity": 0.2, "turn": 2, "available_actions": [0, 1, 2, 3],
"settled_price": 950, "q_value_at_selection": 0.1,
}
err = await DB_SESSION_MNG.execute_lambda_run([DBType.MAIN.value], [lambda s: repo.log_transition(s, data)])
assert err == ErrorType.SUCCESS
err, cnt_a = await _read(lambda s: repo.count_experience(s))
assert cnt_a == 1
# 위조한 company_id 로는 조회되지 않는다
err, cnt_atk = await _read(lambda s: LearningRepository("ATTACKER").count_experience(s))
assert cnt_atk == 0
# 신규 컬럼 값 확인
from sqlalchemy import select
from common.database.model.models import ExperienceLog
err, rows = await _read(lambda s: DB_SESSION_MNG.execute(s, select(ExperienceLog).where(ExperienceLog.company_id == COMPANY_A)))
log = rows[0]
assert log.propensity == 0.2
assert log.turn == 2
assert log.available_actions == [0, 1, 2, 3]
assert log.settled_price == 950
assert log.snapshot["round_number"] == 2
@pytest.mark.asyncio
async def test_reset_all_isolates_tenants(db_engine):
"""파괴 테스트: company A 의 reset_all 이 company B 데이터를 건드리면 안 된다."""
repo_a = LearningRepository(COMPANY_A)
repo_b = LearningRepository(COMPANY_B)
base = dict(state_index=10, action_id=1, card_id="x")
for repo, cid in ((repo_a, COMPANY_A), (repo_b, COMPANY_B)):
for _ in range(3):
await DB_SESSION_MNG.execute_lambda_run([DBType.MAIN.value], [lambda s, r=repo: r.log_transition(s, dict(base))])
err, ca = await _read(lambda s: repo_a.count_experience(s))
err, cb = await _read(lambda s: repo_b.count_experience(s))
assert ca == 3 and cb == 3
# A 만 리셋
assert await repo_a.reset_all() == ErrorType.SUCCESS
err, ca2 = await _read(lambda s: repo_a.count_experience(s))
err, cb2 = await _read(lambda s: repo_b.count_experience(s))
assert ca2 == 0 # A 의 데이터는 사라짐
assert cb2 == 3 # B 의 데이터는 그대로 (타테넌트 무영향)
@pytest.mark.asyncio
async def test_repo_requires_company_id():
with pytest.raises(ValueError):
LearningRepository("")