multi-agent-system / tests /test_token_usage_router.py
firepenguindisopanda
feat(01-01): consolidate SQLAlchemy models into single models.py
d9aaa76
Raw
History Blame Contribute Delete
3.92 kB
"""Tests for token usage API endpoints."""
import pytest
from fastapi.testclient import TestClient
from sqlalchemy import event, text
from app.core.database import engine
from app.core.models import Base
from app.core.models import (
TokenUsageRecord, # noqa: F401 - registers model with Base
)
from app.main import app
@event.listens_for(engine, "connect")
def set_sqlite_pragma(dbapi_connection, connection_record):
"""Enable foreign keys for SQLite."""
if "sqlite" in str(engine.url):
cursor = dbapi_connection.cursor()
cursor.execute("PRAGMA foreign_keys=ON")
cursor.close()
@pytest.fixture(autouse=True)
def setup_test_db():
"""Ensure tables exist and clean up after each test."""
# For SQLite, we need to manually create the table with correct autoincrement
if "sqlite" in str(engine.url):
with engine.connect() as conn:
conn.execute(text("""
CREATE TABLE IF NOT EXISTS token_usage_records (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER,
project_id INTEGER,
operation VARCHAR(100) NOT NULL,
model VARCHAR(100) NOT NULL,
input_tokens INTEGER NOT NULL,
output_tokens INTEGER NOT NULL,
total_tokens INTEGER NOT NULL,
cost_usd FLOAT,
latency_ms INTEGER,
status VARCHAR(20) NOT NULL DEFAULT 'success',
error_message TEXT,
created_at DATETIME DEFAULT CURRENT_TIMESTAMP NOT NULL
)
"""))
conn.commit()
else:
Base.metadata.create_all(bind=engine)
yield
# Clean up token_usage_records table after each test
with engine.connect() as conn:
conn.execute(text("DELETE FROM token_usage_records"))
conn.commit()
def test_create_token_usage():
"""Test creating a token usage record."""
client = TestClient(app)
response = client.post(
"/api/token-usage/",
json={
"operation": "prd_generation",
"model": "gpt-4",
"input_tokens": 1500,
"output_tokens": 3000,
"total_tokens": 4500,
"cost_usd": 0.15,
"latency_ms": 2500,
},
)
assert response.status_code == 201
data = response.json()
assert data["operation"] == "prd_generation"
assert data["total_tokens"] == 4500
assert data["cost_usd"] == 0.15
def test_list_token_usage():
"""Test listing token usage records."""
client = TestClient(app)
# Create some records first
for _ in range(3):
client.post(
"/api/token-usage/",
json={
"operation": "test_op",
"model": "test-model",
"input_tokens": 100,
"output_tokens": 200,
"total_tokens": 300,
},
)
response = client.get("/api/token-usage/")
assert response.status_code == 200
data = response.json()
assert len(data) == 3
def test_get_token_usage_summary():
"""Test getting aggregated token usage summary."""
client = TestClient(app)
# Create records
for _ in range(5):
client.post(
"/api/token-usage/",
json={
"operation": "test_op",
"model": "test-model",
"input_tokens": 100,
"output_tokens": 200,
"total_tokens": 300,
"cost_usd": 0.01,
"latency_ms": 1000,
},
)
response = client.get("/api/token-usage/summary")
assert response.status_code == 200
data = response.json()
assert data["total_operations"] == 5
assert data["total_tokens"] == 1500
assert data["total_cost"] == 0.05