127 lines
5.4 KiB
Python
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("")
|