"""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