Skip to content

Commit 8a9fc0d

Browse files
authored
Fix pricing cache and add detailed cost breakdown (#69)
1 parent fced43a commit 8a9fc0d

5 files changed

Lines changed: 119 additions & 33 deletions

File tree

app/api/routes/claude_code.py

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414
from pydantic import ValidationError
1515
from sqlalchemy.ext.asyncio import AsyncSession
1616

17-
from app.api.dependencies import get_user_by_api_key
17+
from app.api.dependencies import get_user_by_api_key, get_user_details_by_api_key
1818
from app.api.routes.proxy import _get_allowed_provider_names
1919
from app.api.schemas.anthropic import (
2020
AnthropicErrorResponse,
@@ -98,13 +98,16 @@ async def _log_and_return_error_response(
9898
@router.post("/messages", response_model=None, tags=["Claude Code"], status_code=200)
9999
async def create_message_proxy(
100100
request: Request,
101-
user: User = Depends(get_user_by_api_key),
101+
user_details: dict[str, Any] = Depends(get_user_details_by_api_key),
102102
db: AsyncSession = Depends(get_async_db),
103103
) -> Union[JSONResponse, StreamingResponse]:
104104
"""
105105
Main endpoint for Claude Code message completions, proxied through Forge to providers.
106106
Handles request/response conversions, streaming, and dynamic model selection.
107107
"""
108+
user = user_details["user"]
109+
api_key_id = user_details["api_key_id"]
110+
108111
request_id = str(uuid.uuid4())
109112
request.state.request_id = request_id
110113
request.state.start_time_monotonic = time.monotonic()
@@ -224,7 +227,7 @@ async def create_message_proxy(
224227

225228
try:
226229
# Use Forge's provider service to process the request
227-
provider_service = await ProviderService.async_get_instance(user, db)
230+
provider_service = await ProviderService.async_get_instance(user, db, api_key_id)
228231
allowed_provider_names = await _get_allowed_provider_names(request, db)
229232

230233
# Process request through Forge

app/api/routes/statistic.py

Lines changed: 85 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,11 @@
77
from enum import StrEnum
88
import decimal
99

10-
from app.api.dependencies import get_async_db, get_current_active_user, get_current_active_user_from_clerk
10+
from app.api.dependencies import (
11+
get_async_db,
12+
get_current_active_user,
13+
get_current_active_user_from_clerk,
14+
)
1115
from app.models.user import User
1216
from app.models.usage_tracker import UsageTracker
1317
from app.models.provider_key import ProviderKey
@@ -40,7 +44,11 @@ async def get_usage_realtime(
4044
# Calculate the date 7 days ago
4145
now = datetime.now(UTC)
4246
seven_days_ago = now - timedelta(days=7)
43-
started_at = started_at if started_at is not None and started_at > seven_days_ago else seven_days_ago
47+
started_at = (
48+
started_at
49+
if started_at is not None and started_at > seven_days_ago
50+
else seven_days_ago
51+
)
4452
ended_at = ended_at if ended_at is not None and ended_at < now else None
4553

4654
# Build the query
@@ -51,6 +59,11 @@ async def get_usage_realtime(
5159
ProviderKey.provider_name.label("provider_name"),
5260
UsageTracker.model.label("model_name"),
5361
(UsageTracker.input_tokens + UsageTracker.output_tokens).label("tokens"),
62+
(UsageTracker.input_tokens - UsageTracker.cached_tokens).label(
63+
"input_tokens"
64+
),
65+
UsageTracker.output_tokens.label("output_tokens"),
66+
UsageTracker.cached_tokens.label("cached_tokens"),
5467
UsageTracker.cost.label("cost"),
5568
func.extract(
5669
"epoch", UsageTracker.updated_at - UsageTracker.created_at
@@ -62,13 +75,15 @@ async def get_usage_realtime(
6275
UsageTracker.user_id == current_user.id,
6376
UsageTracker.created_at >= started_at,
6477
ended_at is None or UsageTracker.created_at <= ended_at,
65-
forge_key is None or or_(
78+
forge_key is None
79+
or or_(
6680
ForgeApiKey.key.ilike(f"%{forge_key}%"),
67-
ForgeApiKey.name.ilike(f"%{forge_key}%")
81+
ForgeApiKey.name.ilike(f"%{forge_key}%"),
6882
),
69-
provider_name is None or ProviderKey.provider_name.ilike(f"%{provider_name}%"),
83+
provider_name is None
84+
or ProviderKey.provider_name.ilike(f"%{provider_name}%"),
7085
model_name is None or UsageTracker.model.ilike(f"%{model_name}%"),
71-
UsageTracker.updated_at.is_not(None)
86+
UsageTracker.updated_at.is_not(None),
7287
)
7388
.order_by(desc(UsageTracker.created_at))
7489
.offset(offset)
@@ -89,16 +104,18 @@ async def get_usage_realtime(
89104
"provider_name": row.provider_name,
90105
"model_name": row.model_name,
91106
"tokens": row.tokens,
107+
"input_tokens": row.input_tokens,
108+
"output_tokens": row.output_tokens,
109+
"cached_tokens": row.cached_tokens,
92110
"cost": decimal.Decimal(row.cost).normalize(),
93111
"duration": round(float(row.duration), 2)
94112
if row.duration is not None
95113
else 0.0,
96114
}
97115
)
98-
print(usage_stats)
99-
100116
return [UsageRealtimeResponse(**usage_stat) for usage_stat in usage_stats]
101117

118+
102119
@router.get("/usage/realtime/clerk", response_model=list[UsageRealtimeResponse])
103120
async def get_usage_realtime_clerk(
104121
current_user: User = Depends(get_current_active_user_from_clerk),
@@ -111,7 +128,17 @@ async def get_usage_realtime_clerk(
111128
started_at: datetime = Query(None),
112129
ended_at: datetime = Query(None),
113130
):
114-
return await get_usage_realtime(current_user, db, offset, limit, forge_key, provider_name, model_name, started_at, ended_at)
131+
return await get_usage_realtime(
132+
current_user,
133+
db,
134+
offset,
135+
limit,
136+
forge_key,
137+
provider_name,
138+
model_name,
139+
started_at,
140+
ended_at,
141+
)
115142

116143

117144
class UsageSummaryTimeSpan(StrEnum):
@@ -152,16 +179,21 @@ async def get_usage_summary(
152179
func.sum(UsageTracker.input_tokens + UsageTracker.output_tokens).label(
153180
"tokens"
154181
),
182+
func.sum(UsageTracker.input_tokens - UsageTracker.cached_tokens).label(
183+
"input_tokens"
184+
),
185+
func.sum(UsageTracker.output_tokens).label("output_tokens"),
186+
func.sum(UsageTracker.cached_tokens).label("cached_tokens"),
155187
func.sum(UsageTracker.cost).label("cost"),
156188
)
157189
.join(ForgeApiKey, UsageTracker.forge_key_id == ForgeApiKey.id)
158190
.where(
159191
UsageTracker.user_id == current_user.id,
160192
UsageTracker.created_at >= start_time,
161-
UsageTracker.updated_at.is_not(None)
193+
UsageTracker.updated_at.is_not(None),
162194
)
163195
.group_by(time_group, ForgeApiKey.name, ForgeApiKey.key)
164-
.order_by(time_group, desc("tokens"), "forge_key")
196+
.order_by(time_group, desc("cost"), "forge_key")
165197
)
166198

167199
# Execute the query
@@ -171,19 +203,41 @@ async def get_usage_summary(
171203
data_points = dict()
172204
for row in rows:
173205
if row.time_point not in data_points:
174-
data_points[row.time_point] = {"breakdown": [], "total_tokens": 0, "total_cost": 0}
206+
data_points[row.time_point] = {
207+
"breakdown": [],
208+
"total_tokens": 0,
209+
"total_cost": 0,
210+
"total_input_tokens": 0,
211+
"total_output_tokens": 0,
212+
"total_cached_tokens": 0,
213+
}
175214
data_points[row.time_point]["breakdown"].append(
176-
{"forge_key": row.forge_key, "tokens": row.tokens, "cost": decimal.Decimal(row.cost).normalize()}
215+
{
216+
"forge_key": row.forge_key,
217+
"tokens": row.tokens,
218+
"cost": decimal.Decimal(row.cost).normalize(),
219+
"input_tokens": row.input_tokens,
220+
"output_tokens": row.output_tokens,
221+
"cached_tokens": row.cached_tokens,
222+
}
177223
)
178224
data_points[row.time_point]["total_tokens"] += row.tokens
179-
data_points[row.time_point]["total_cost"] += decimal.Decimal(row.cost).normalize()
225+
data_points[row.time_point]["total_cost"] += decimal.Decimal(
226+
row.cost
227+
).normalize()
228+
data_points[row.time_point]["total_input_tokens"] += row.input_tokens
229+
data_points[row.time_point]["total_output_tokens"] += row.output_tokens
230+
data_points[row.time_point]["total_cached_tokens"] += row.cached_tokens
180231

181232
return [
182233
UsageSummaryResponse(
183234
time_point=time_point,
184235
breakdown=data_point["breakdown"],
185236
total_tokens=data_point["total_tokens"],
186237
total_cost=data_point["total_cost"],
238+
total_input_tokens=data_point["total_input_tokens"],
239+
total_output_tokens=data_point["total_output_tokens"],
240+
total_cached_tokens=data_point["total_cached_tokens"],
187241
)
188242
for time_point, data_point in data_points.items()
189243
]
@@ -231,6 +285,11 @@ async def get_forge_keys_usage(
231285
func.sum(UsageTracker.input_tokens + UsageTracker.output_tokens).label(
232286
"tokens"
233287
),
288+
func.sum(UsageTracker.input_tokens - UsageTracker.cached_tokens).label(
289+
"input_tokens"
290+
),
291+
func.sum(UsageTracker.output_tokens).label("output_tokens"),
292+
func.sum(UsageTracker.cached_tokens).label("cached_tokens"),
234293
func.sum(UsageTracker.cost).label("cost"),
235294
)
236295
.join(ForgeApiKey, UsageTracker.forge_key_id == ForgeApiKey.id)
@@ -240,19 +299,28 @@ async def get_forge_keys_usage(
240299
UsageTracker.updated_at.is_not(None),
241300
)
242301
.group_by(ForgeApiKey.name, ForgeApiKey.key)
243-
.order_by(desc("tokens"), "forge_key")
302+
.order_by(desc("cost"), "forge_key")
244303
)
245304

246305
result = await db.execute(query)
247306
rows = result.fetchall()
248307

249308
return [
250-
ForgeKeysUsageSummaryResponse(forge_key=row.forge_key, tokens=row.tokens, cost=decimal.Decimal(row.cost).normalize())
309+
ForgeKeysUsageSummaryResponse(
310+
forge_key=row.forge_key,
311+
tokens=row.tokens,
312+
cost=decimal.Decimal(row.cost).normalize(),
313+
input_tokens=row.input_tokens,
314+
output_tokens=row.output_tokens,
315+
cached_tokens=row.cached_tokens,
316+
)
251317
for row in rows
252318
]
253319

254320

255-
@router.get("/forge-keys/usage/clerk", response_model=list[ForgeKeysUsageSummaryResponse])
321+
@router.get(
322+
"/forge-keys/usage/clerk", response_model=list[ForgeKeysUsageSummaryResponse]
323+
)
256324
async def get_forge_keys_usage_clerk(
257325
current_user: User = Depends(get_current_active_user_from_clerk),
258326
db: AsyncSession = Depends(get_async_db),

app/api/schemas/statistic.py

Lines changed: 20 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -13,11 +13,14 @@ def mask_forge_name_or_key(v: str) -> str:
1313
return v
1414

1515
class UsageRealtimeResponse(BaseModel):
16-
timestamp: datetime
16+
timestamp: datetime | str
1717
forge_key: str
1818
provider_name: str
1919
model_name: str
2020
tokens: int
21+
input_tokens: int
22+
output_tokens: int
23+
cached_tokens: int
2124
duration: float
2225
cost: decimal.Decimal
2326

@@ -28,14 +31,19 @@ def mask_forge_key(cls, v: str) -> str:
2831

2932
@field_validator('timestamp')
3033
@classmethod
31-
def convert_timestamp_to_iso(cls, v: datetime) -> str:
34+
def convert_timestamp_to_iso(cls, v: datetime | str) -> str:
35+
if isinstance(v, str):
36+
return v
3237
return v.isoformat()
3338

3439

3540
class UsageSummaryBreakdown(BaseModel):
3641
forge_key: str
3742
tokens: int
3843
cost: decimal.Decimal
44+
input_tokens: int
45+
output_tokens: int
46+
cached_tokens: int
3947

4048
@field_validator('forge_key')
4149
@classmethod
@@ -44,21 +52,29 @@ def mask_forge_key(cls, v: str) -> str:
4452

4553

4654
class UsageSummaryResponse(BaseModel):
47-
time_point: datetime
55+
time_point: datetime | str
4856
breakdown: list[UsageSummaryBreakdown]
4957
total_tokens: int
5058
total_cost: decimal.Decimal
59+
total_input_tokens: int
60+
total_output_tokens: int
61+
total_cached_tokens: int
5162

5263
@field_validator('time_point')
5364
@classmethod
54-
def convert_timestamp_to_iso(cls, v: datetime) -> str:
65+
def convert_timestamp_to_iso(cls, v: datetime | str) -> str:
66+
if isinstance(v, str):
67+
return v
5568
return v.isoformat()
5669

5770

5871
class ForgeKeysUsageSummaryResponse(BaseModel):
5972
forge_key: str
6073
tokens: int
6174
cost: decimal.Decimal
75+
input_tokens: int
76+
output_tokens: int
77+
cached_tokens: int
6278

6379
@field_validator('forge_key')
6480
@classmethod

app/services/pricing_service.py

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -93,23 +93,23 @@ async def _fetch_pricing_with_smart_caching(
9393
exact_pricing = await async_provider_service_cache.get(exact_cache_key)
9494

9595
if exact_pricing and PricingService._is_pricing_valid_for_date(exact_pricing, calculation_date):
96-
logger.debug(f"Hot cache hit for {provider_name}/{model_name}")
96+
logger.debug(f"Exact match cache hit for {provider_name}/{model_name}")
9797
return {**exact_pricing, 'source': 'exact_match'}
9898

9999
# Level 2: Try provider fallback cache (warm cache)
100-
provider_cache_key = f"pricing:provider_fallback:{provider_name}"
100+
provider_cache_key = f"pricing:provider_fallback:{provider_name}:{model_name}"
101101
provider_fallback = await async_provider_service_cache.get(provider_cache_key)
102102

103103
if provider_fallback and PricingService._is_pricing_valid_for_date(provider_fallback, calculation_date):
104-
logger.debug(f"Warm cache hit for provider {provider_name}")
104+
logger.debug(f"Provider fallback cache hit for {provider_name}")
105105
return {**provider_fallback, 'source': 'fallback_provider'}
106106

107107
# Level 3: Try global fallback cache (warm cache)
108-
global_cache_key = f"pricing:global_fallback"
108+
global_cache_key = f"pricing:global_fallback:{provider_name}:{model_name}"
109109
global_fallback = await async_provider_service_cache.get(global_cache_key)
110110

111111
if global_fallback and PricingService._is_pricing_valid_for_date(global_fallback, calculation_date):
112-
logger.debug("Warm cache hit for global fallback")
112+
logger.debug("Global fallback cache hit")
113113
return {**global_fallback, 'source': 'fallback_global'}
114114

115115
# Cache miss - hit database (this should be rare)
@@ -149,7 +149,7 @@ async def _fetch_from_database_with_caching(
149149

150150
if provider_fallback:
151151
# Cache provider fallback (warm cache)
152-
cache_key = f"pricing:provider_fallback:{provider_name}"
152+
cache_key = f"pricing:provider_fallback:{provider_name}:{model_name}"
153153
await async_provider_service_cache.set(
154154
cache_key, provider_fallback, ttl=PricingService.FALLBACK_CACHE_TTL
155155
)
@@ -163,7 +163,7 @@ async def _fetch_from_database_with_caching(
163163

164164
if global_fallback:
165165
# Cache global fallback (warm cache)
166-
cache_key = f"pricing:global_fallback"
166+
cache_key = f"pricing:global_fallback:{provider_name}:{model_name}"
167167
await async_provider_service_cache.set(
168168
cache_key, global_fallback, ttl=PricingService.FALLBACK_CACHE_TTL
169169
)

app/services/provider_service.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -598,8 +598,7 @@ async def process_request(
598598
endpoint=endpoint,
599599
)
600600
else:
601-
# TODO: this shouldn't happen, but we handle it gracefully as we don't want to break the flow
602-
# Dive deeper into this if it ever happens
601+
# For api like list models, we don't have usage tracking
603602
logger.info(
604603
f"api_key_id: {self.api_key_id}, provider_key_id: {provider_key_id}"
605604
)

0 commit comments

Comments
 (0)