mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat: 视频计费增强与影子计费系统
This commit is contained in:
301
tests/services/billing/test_default_rules.py
Normal file
301
tests/services/billing/test_default_rules.py
Normal file
@@ -0,0 +1,301 @@
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from src.models.database import GlobalModel, Model
|
||||
from src.services.billing.default_rules import DefaultBillingRuleGenerator
|
||||
from src.services.billing.formula_engine import FormulaEngine
|
||||
from src.services.billing.rule_service import BillingRuleService
|
||||
|
||||
|
||||
class TestDefaultBillingRuleGenerator:
|
||||
def test_default_rule_basic_chat_cost(self) -> None:
|
||||
global_model = GlobalModel(
|
||||
name="test-model",
|
||||
display_name="Test Model",
|
||||
is_active=True,
|
||||
default_price_per_request=0.01,
|
||||
default_tiered_pricing={
|
||||
"tiers": [
|
||||
{
|
||||
"up_to": None,
|
||||
"input_price_per_1m": 3.0,
|
||||
"output_price_per_1m": 15.0,
|
||||
# cache prices intentionally omitted (legacy derives from input price)
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
rule = DefaultBillingRuleGenerator.generate_for_model(
|
||||
global_model=global_model,
|
||||
model=None,
|
||||
task_type="chat",
|
||||
)
|
||||
|
||||
engine = FormulaEngine()
|
||||
result = engine.evaluate(
|
||||
expression=rule.expression,
|
||||
variables=rule.variables,
|
||||
dimensions={
|
||||
"input_tokens": 1000,
|
||||
"output_tokens": 500,
|
||||
"cache_creation_tokens": 200,
|
||||
"cache_read_tokens": 300,
|
||||
"request_count": 1,
|
||||
# tier key
|
||||
"total_input_context": 1000 + 300,
|
||||
},
|
||||
dimension_mappings=rule.dimension_mappings,
|
||||
strict_mode=True,
|
||||
)
|
||||
|
||||
assert result.status == "complete"
|
||||
assert abs(float(result.cost) - 0.02134) < 1e-9
|
||||
|
||||
def test_default_rule_tiered_pricing_uses_total_input_context(self) -> None:
|
||||
global_model = GlobalModel(
|
||||
name="tiered-model",
|
||||
display_name="Tiered Model",
|
||||
is_active=True,
|
||||
default_price_per_request=None,
|
||||
default_tiered_pricing={
|
||||
"tiers": [
|
||||
{"up_to": 200000, "input_price_per_1m": 3.0, "output_price_per_1m": 15.0},
|
||||
{"up_to": None, "input_price_per_1m": 1.5, "output_price_per_1m": 7.5},
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
rule = DefaultBillingRuleGenerator.generate_for_model(
|
||||
global_model=global_model,
|
||||
model=None,
|
||||
task_type="chat",
|
||||
)
|
||||
|
||||
engine = FormulaEngine()
|
||||
result = engine.evaluate(
|
||||
expression=rule.expression,
|
||||
variables=rule.variables,
|
||||
dimensions={
|
||||
"input_tokens": 250000,
|
||||
"output_tokens": 10000,
|
||||
"cache_creation_tokens": 0,
|
||||
"cache_read_tokens": 0,
|
||||
"request_count": 1,
|
||||
"total_input_context": 250000,
|
||||
},
|
||||
dimension_mappings=rule.dimension_mappings,
|
||||
strict_mode=True,
|
||||
)
|
||||
|
||||
# 250000 * 1.5 / 1M = 0.375
|
||||
# 10000 * 7.5 / 1M = 0.075
|
||||
assert result.status == "complete"
|
||||
assert abs(float(result.cost) - 0.45) < 1e-9
|
||||
|
||||
def test_default_rule_cache_ttl_pricing_overrides_cache_read_price(self) -> None:
|
||||
global_model = GlobalModel(
|
||||
name="ttl-model",
|
||||
display_name="TTL Model",
|
||||
is_active=True,
|
||||
default_price_per_request=0.0,
|
||||
default_tiered_pricing={
|
||||
"tiers": [
|
||||
{
|
||||
"up_to": None,
|
||||
"input_price_per_1m": 3.0,
|
||||
"output_price_per_1m": 15.0,
|
||||
"cache_read_price_per_1m": 0.3,
|
||||
"cache_ttl_pricing": [
|
||||
{"ttl_minutes": 5, "cache_read_price_per_1m": 0.3},
|
||||
{"ttl_minutes": 60, "cache_read_price_per_1m": 0.5},
|
||||
],
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
rule = DefaultBillingRuleGenerator.generate_for_model(
|
||||
global_model=global_model,
|
||||
model=None,
|
||||
task_type="chat",
|
||||
)
|
||||
|
||||
engine = FormulaEngine()
|
||||
result = engine.evaluate(
|
||||
expression=rule.expression,
|
||||
variables=rule.variables,
|
||||
dimensions={
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
"cache_creation_tokens": 0,
|
||||
"cache_read_tokens": 1000,
|
||||
"cache_ttl_minutes": 60,
|
||||
"request_count": 1,
|
||||
"total_input_context": 0 + 1000,
|
||||
},
|
||||
dimension_mappings=rule.dimension_mappings,
|
||||
strict_mode=True,
|
||||
)
|
||||
|
||||
# TTL=60 should use cache_read_price_per_1m=0.5
|
||||
# 1000 * 0.5 / 1M = 0.0005
|
||||
assert result.status == "complete"
|
||||
assert abs(float(result.cost) - 0.0005) < 1e-9
|
||||
|
||||
|
||||
class TestBillingRuleServiceDefaultFallback:
|
||||
def test_find_rule_returns_default_for_chat_when_no_db_rule(self) -> None:
|
||||
from src.services.billing.cache import BillingCache
|
||||
|
||||
BillingCache.invalidate_all()
|
||||
|
||||
global_model = GlobalModel(
|
||||
id="gm-1",
|
||||
name="test-model",
|
||||
display_name="Test Model",
|
||||
is_active=True,
|
||||
default_price_per_request=0.0,
|
||||
default_tiered_pricing={
|
||||
"tiers": [{"up_to": None, "input_price_per_1m": 3.0, "output_price_per_1m": 15.0}]
|
||||
},
|
||||
)
|
||||
|
||||
model_obj = Model(
|
||||
id="m-1",
|
||||
provider_id="p-1",
|
||||
global_model_id="gm-1",
|
||||
provider_model_name="provider-test-model",
|
||||
is_active=True,
|
||||
tiered_pricing=None,
|
||||
price_per_request=None,
|
||||
)
|
||||
model_obj.global_model = global_model
|
||||
|
||||
# Build a mock Session with deterministic query().filter().first() chain.
|
||||
q_global = MagicMock()
|
||||
q_global.filter.return_value.first.return_value = global_model
|
||||
|
||||
q_model = MagicMock()
|
||||
q_model.filter.return_value.first.return_value = model_obj
|
||||
|
||||
q_rule_model = MagicMock()
|
||||
q_rule_model.filter.return_value.first.return_value = None
|
||||
|
||||
q_rule_global = MagicMock()
|
||||
q_rule_global.filter.return_value.first.return_value = None
|
||||
|
||||
db = MagicMock()
|
||||
db.query.side_effect = [q_global, q_model, q_rule_model, q_rule_global]
|
||||
|
||||
lookup = BillingRuleService.find_rule(
|
||||
db,
|
||||
provider_id="p-1",
|
||||
model_name="test-model",
|
||||
task_type="chat",
|
||||
)
|
||||
assert lookup is not None
|
||||
assert lookup.scope == "default"
|
||||
assert lookup.rule.id == "__default__"
|
||||
assert lookup.effective_task_type == "chat"
|
||||
|
||||
# Cached: second call should not touch db.query again.
|
||||
db.query.reset_mock()
|
||||
lookup2 = BillingRuleService.find_rule(
|
||||
db,
|
||||
provider_id="p-1",
|
||||
model_name="test-model",
|
||||
task_type="chat",
|
||||
)
|
||||
assert lookup2 is not None
|
||||
assert lookup2.scope == "default"
|
||||
assert db.query.call_count == 0
|
||||
|
||||
def test_find_rule_returns_default_for_video_when_require_rule_false(self) -> None:
|
||||
from src.config.settings import config
|
||||
from src.services.billing.cache import BillingCache
|
||||
|
||||
BillingCache.invalidate_all()
|
||||
old_require = config.billing_require_rule
|
||||
config.billing_require_rule = False
|
||||
|
||||
try:
|
||||
global_model = GlobalModel(
|
||||
id="gm-2",
|
||||
name="video-model",
|
||||
display_name="Video Model",
|
||||
is_active=True,
|
||||
default_price_per_request=0.0,
|
||||
default_tiered_pricing={
|
||||
"tiers": [
|
||||
{"up_to": None, "input_price_per_1m": 3.0, "output_price_per_1m": 15.0}
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
q_global = MagicMock()
|
||||
q_global.filter.return_value.first.return_value = global_model
|
||||
|
||||
q_rule_global = MagicMock()
|
||||
q_rule_global.filter.return_value.first.return_value = None
|
||||
|
||||
db = MagicMock()
|
||||
# provider_id omitted -> only GlobalModel query + global BillingRule query
|
||||
db.query.side_effect = [q_global, q_rule_global]
|
||||
|
||||
lookup = BillingRuleService.find_rule(
|
||||
db,
|
||||
provider_id=None,
|
||||
model_name="video-model",
|
||||
task_type="video",
|
||||
)
|
||||
assert lookup is not None
|
||||
assert lookup.scope == "default"
|
||||
assert lookup.rule.id == "__default__"
|
||||
assert lookup.effective_task_type == "video"
|
||||
finally:
|
||||
config.billing_require_rule = old_require
|
||||
|
||||
def test_find_rule_returns_template_for_video_when_require_rule_true(self) -> None:
|
||||
from src.config.settings import config
|
||||
from src.services.billing.cache import BillingCache
|
||||
|
||||
BillingCache.invalidate_all()
|
||||
old_require = config.billing_require_rule
|
||||
config.billing_require_rule = True
|
||||
|
||||
try:
|
||||
global_model = GlobalModel(
|
||||
id="gm-3",
|
||||
name="video-model-2",
|
||||
display_name="Video Model 2",
|
||||
is_active=True,
|
||||
default_price_per_request=0.0,
|
||||
default_tiered_pricing={
|
||||
"tiers": [
|
||||
{"up_to": None, "input_price_per_1m": 3.0, "output_price_per_1m": 15.0}
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
q_global = MagicMock()
|
||||
q_global.filter.return_value.first.return_value = global_model
|
||||
|
||||
q_rule_global = MagicMock()
|
||||
q_rule_global.filter.return_value.first.return_value = None
|
||||
|
||||
db = MagicMock()
|
||||
db.query.side_effect = [q_global, q_rule_global]
|
||||
|
||||
lookup = BillingRuleService.find_rule(
|
||||
db,
|
||||
provider_id=None,
|
||||
model_name="video-model-2",
|
||||
task_type="video",
|
||||
)
|
||||
assert lookup is not None
|
||||
assert lookup.scope == "default"
|
||||
assert lookup.rule.id == "__default__"
|
||||
# Universal Billing Rule is now used for all task types
|
||||
assert lookup.rule.name == "Universal Billing Rule"
|
||||
finally:
|
||||
config.billing_require_rule = old_require
|
||||
@@ -34,7 +34,7 @@ class TestDimensionCollectorRuntime:
|
||||
),
|
||||
]
|
||||
dims = runtime.collect(
|
||||
collectors=collectors,
|
||||
collectors=collectors, # type: ignore[arg-type]
|
||||
inp=DimensionCollectInput(
|
||||
response={"usageMetadata": {"promptTokenCount": 123}},
|
||||
),
|
||||
@@ -57,7 +57,7 @@ class TestDimensionCollectorRuntime:
|
||||
)
|
||||
]
|
||||
dims = runtime.collect(
|
||||
collectors=collectors,
|
||||
collectors=collectors, # type: ignore[arg-type]
|
||||
inp=DimensionCollectInput(metadata={"result": {"file_size_bytes": 1048576}}),
|
||||
)
|
||||
assert abs(dims["file_size_mb"] - 1.0) < 1e-9
|
||||
@@ -98,7 +98,7 @@ class TestDimensionCollectorRuntime:
|
||||
),
|
||||
]
|
||||
dims = runtime.collect(
|
||||
collectors=collectors,
|
||||
collectors=collectors, # type: ignore[arg-type]
|
||||
inp=DimensionCollectInput(
|
||||
request={"usage": {"input_tokens": 100, "cache_read_tokens": 20}}
|
||||
),
|
||||
@@ -110,52 +110,13 @@ class TestDimensionCollectorRuntime:
|
||||
|
||||
class TestDimensionCollectorService:
|
||||
def test_video_fallback_merges_base_collectors(self) -> None:
|
||||
from src.services.billing.cache import BillingCache
|
||||
|
||||
BillingCache.invalidate_all()
|
||||
|
||||
# code-only: openai:video should fall back to openai:chat video collectors shipped in code
|
||||
db = MagicMock()
|
||||
|
||||
video_collectors = [
|
||||
DimensionCollector(
|
||||
api_format="openai:video",
|
||||
task_type="video",
|
||||
dimension_name="duration_seconds",
|
||||
source_type="metadata",
|
||||
source_path="task.duration_seconds",
|
||||
value_type="int",
|
||||
priority=0,
|
||||
is_enabled=True,
|
||||
)
|
||||
]
|
||||
base_collectors = [
|
||||
# Should be kept (dimension not present in video_collectors)
|
||||
DimensionCollector(
|
||||
api_format="openai:chat",
|
||||
task_type="video",
|
||||
dimension_name="resolution",
|
||||
source_type="metadata",
|
||||
source_path="task.resolution",
|
||||
value_type="string",
|
||||
priority=0,
|
||||
is_enabled=True,
|
||||
),
|
||||
# Should be ignored (dimension already present in video_collectors)
|
||||
DimensionCollector(
|
||||
api_format="openai:chat",
|
||||
task_type="video",
|
||||
dimension_name="duration_seconds",
|
||||
source_type="metadata",
|
||||
source_path="task.duration_seconds",
|
||||
value_type="int",
|
||||
priority=0,
|
||||
is_enabled=True,
|
||||
),
|
||||
]
|
||||
|
||||
q1 = MagicMock()
|
||||
q1.filter.return_value.all.return_value = video_collectors
|
||||
q2 = MagicMock()
|
||||
q2.filter.return_value.all.return_value = base_collectors
|
||||
db.query.side_effect = [q1, q2]
|
||||
|
||||
svc = DimensionCollectorService(db)
|
||||
result = svc.list_enabled_collectors(api_format="openai:video", task_type="video")
|
||||
|
||||
assert [c.dimension_name for c in result] == ["duration_seconds", "resolution"]
|
||||
assert "video_size_bytes" in [c.dimension_name for c in result]
|
||||
|
||||
@@ -54,7 +54,7 @@ class TestFormulaEngine:
|
||||
)
|
||||
assert result.status == "complete"
|
||||
assert result.missing_required == []
|
||||
assert abs(result.cost - 0.25) < 1e-9
|
||||
assert abs(float(result.cost) - 0.25) < 1e-9
|
||||
|
||||
def test_required_dimension_missing_non_strict(self) -> None:
|
||||
engine = FormulaEngine()
|
||||
@@ -72,7 +72,7 @@ class TestFormulaEngine:
|
||||
strict_mode=False,
|
||||
)
|
||||
assert result.status == "incomplete"
|
||||
assert result.cost == 0.0
|
||||
assert float(result.cost) == 0.0
|
||||
assert result.missing_required == ["duration_seconds"]
|
||||
|
||||
def test_required_dimension_missing_strict_raises(self) -> None:
|
||||
@@ -127,7 +127,7 @@ class TestFormulaEngine:
|
||||
},
|
||||
)
|
||||
assert result.status == "complete"
|
||||
assert result.cost == 0.0
|
||||
assert float(result.cost) == 0.0
|
||||
|
||||
def test_tiered_mapping(self) -> None:
|
||||
engine = FormulaEngine()
|
||||
@@ -148,7 +148,7 @@ class TestFormulaEngine:
|
||||
},
|
||||
)
|
||||
assert result.status == "complete"
|
||||
assert result.cost == 3.0
|
||||
assert float(result.cost) == 3.0
|
||||
|
||||
result = engine.evaluate(
|
||||
expression="input_price",
|
||||
@@ -166,4 +166,4 @@ class TestFormulaEngine:
|
||||
},
|
||||
)
|
||||
assert result.status == "complete"
|
||||
assert result.cost == 1.5
|
||||
assert float(result.cost) == 1.5
|
||||
|
||||
112
tests/services/billing/test_shadow_billing.py
Normal file
112
tests/services/billing/test_shadow_billing.py
Normal file
@@ -0,0 +1,112 @@
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from src.config.settings import config
|
||||
from src.services.billing.schema import BillingSnapshot, CostResult
|
||||
from src.services.billing.shadow import CostBreakdown, ShadowBillingService
|
||||
|
||||
|
||||
class TestShadowBillingServiceModeResolution:
|
||||
def test_get_engine_mode_exact_and_wildcard(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(config, "billing_engine", "legacy", raising=False)
|
||||
monkeypatch.setattr(
|
||||
config,
|
||||
"billing_engine_overrides",
|
||||
'{"anthropic/*": "shadow", "openai/gpt-4o": "new"}',
|
||||
raising=False,
|
||||
)
|
||||
|
||||
svc = ShadowBillingService(MagicMock())
|
||||
assert svc.get_engine_mode("openai", "gpt-4o") == "new"
|
||||
assert svc.get_engine_mode("anthropic", "claude-3-5-sonnet") == "shadow"
|
||||
assert svc.get_engine_mode("other", "x") == "legacy"
|
||||
|
||||
|
||||
class TestShadowBillingServiceExecution:
|
||||
def test_legacy_mode_skips_new_engine(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(config, "billing_engine", "legacy", raising=False)
|
||||
monkeypatch.setattr(config, "billing_engine_overrides", "{}", raising=False)
|
||||
|
||||
svc = ShadowBillingService(MagicMock())
|
||||
# Guard: if new engine calculate gets called, fail.
|
||||
svc._new_billing = MagicMock()
|
||||
svc._new_billing.calculate.side_effect = AssertionError(
|
||||
"new engine should not run in legacy mode"
|
||||
)
|
||||
|
||||
legacy_truth = CostBreakdown(
|
||||
input_cost=0.1,
|
||||
output_cost=0.2,
|
||||
cache_creation_cost=0.0,
|
||||
cache_read_cost=0.0,
|
||||
request_cost=0.0,
|
||||
total_cost=0.3,
|
||||
)
|
||||
|
||||
res = svc.calculate_with_shadow(
|
||||
provider="openai",
|
||||
provider_id="p-1",
|
||||
model="gpt-4o",
|
||||
task_type="chat",
|
||||
api_format="openai:chat",
|
||||
input_tokens=1,
|
||||
output_tokens=1,
|
||||
legacy_truth=legacy_truth,
|
||||
is_failed_request=False,
|
||||
)
|
||||
|
||||
assert res.engine_mode == "legacy"
|
||||
assert res.truth_engine == "legacy"
|
||||
assert res.shadow_snapshot is None
|
||||
assert res.truth_breakdown.total_cost == 0.3
|
||||
|
||||
def test_shadow_mode_returns_snapshot_and_keeps_legacy_truth(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setattr(config, "billing_engine", "shadow", raising=False)
|
||||
monkeypatch.setattr(config, "billing_engine_overrides", "{}", raising=False)
|
||||
monkeypatch.setattr(config, "billing_diff_threshold_usd", 0.0001, raising=False)
|
||||
|
||||
svc = ShadowBillingService(MagicMock())
|
||||
|
||||
# Stub new engine output
|
||||
snapshot = BillingSnapshot(
|
||||
resolved_dimensions={"input_tokens": 1},
|
||||
resolved_variables={"input_price_per_1m": "3.0"},
|
||||
cost_breakdown={"input_cost": 0.003},
|
||||
total_cost=0.003,
|
||||
status="complete",
|
||||
calculated_at="2026-02-02T00:00:00Z",
|
||||
)
|
||||
svc._new_billing = MagicMock()
|
||||
svc._new_billing.calculate.return_value = CostResult(
|
||||
cost=0.003, status="complete", snapshot=snapshot
|
||||
)
|
||||
|
||||
legacy_truth = CostBreakdown(
|
||||
input_cost=0.004,
|
||||
output_cost=0.0,
|
||||
cache_creation_cost=0.0,
|
||||
cache_read_cost=0.0,
|
||||
request_cost=0.0,
|
||||
total_cost=0.004,
|
||||
)
|
||||
|
||||
res = svc.calculate_with_shadow(
|
||||
provider="openai",
|
||||
provider_id="p-1",
|
||||
model="gpt-4o",
|
||||
task_type="chat",
|
||||
api_format="openai:chat",
|
||||
input_tokens=1,
|
||||
output_tokens=0,
|
||||
legacy_truth=legacy_truth,
|
||||
is_failed_request=False,
|
||||
)
|
||||
|
||||
assert res.engine_mode == "shadow"
|
||||
assert res.truth_engine == "legacy"
|
||||
assert res.shadow_snapshot is not None
|
||||
assert res.truth_breakdown.total_cost == 0.004
|
||||
assert "diff_usd" in res.comparison
|
||||
Reference in New Issue
Block a user