model-router-cli 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
app/api/routes.py ADDED
@@ -0,0 +1,589 @@
1
+ import uuid
2
+ import time
3
+ import asyncio
4
+ import csv
5
+ import io
6
+ from typing import List, Optional, Dict, Any
7
+ from fastapi import APIRouter, Depends, HTTPException, Query
8
+ from fastapi.responses import JSONResponse, StreamingResponse
9
+ from pydantic import BaseModel, Field
10
+ from sqlalchemy.ext.asyncio import AsyncSession
11
+ from sqlalchemy import select, desc
12
+
13
+ from app.storage.database import get_db
14
+ from app.storage.models import (
15
+ ModelRecord,
16
+ RoutingPolicyRecord,
17
+ RoutingRuleRecord,
18
+ RequestRecord,
19
+ RoutingDecisionRecord,
20
+ ResponseRecord,
21
+ FeedbackRecord,
22
+ BudgetRecord,
23
+ ExperimentRecord,
24
+ )
25
+ from app.models.schemas import (
26
+ ModelMetadata,
27
+ RequestAnalysis,
28
+ RoutingDecision,
29
+ ProviderResponse,
30
+ ModelTier,
31
+ )
32
+ from app.analyzer.analyzer import analyze_request
33
+ from app.router.engine import route_request
34
+ from app.fallback.handler import execute_with_fallback
35
+ from app.budgets.manager import check_budget_threshold
36
+ from app.observability.events import log_router_event
37
+ from app.analytics.service import get_system_analytics, calculate_cost_savings
38
+ from app.providers.registry import provider_registry
39
+ from app.config.settings import get_settings
40
+
41
+ router = APIRouter()
42
+ settings = get_settings()
43
+
44
+
45
+ class RouteOnlyRequest(BaseModel):
46
+ prompt: str = Field(..., min_length=1, description="Prompt text to analyze and route")
47
+ policy: Optional[str] = None
48
+ analyzer_mode: Optional[str] = None
49
+
50
+
51
+ class GenerateRequest(BaseModel):
52
+ prompt: str = Field(..., min_length=1)
53
+ system_prompt: Optional[str] = None
54
+ policy: Optional[str] = None
55
+ analyzer_mode: Optional[str] = None
56
+ force_model: Optional[str] = None
57
+ temperature: float = 0.7
58
+ max_retries: int = 1
59
+
60
+
61
+ class FeedbackSubmitRequest(BaseModel):
62
+ request_id: str
63
+ rating: int = Field(..., description="1 for up, -1 for down")
64
+ comment: Optional[str] = None
65
+
66
+
67
+ class ModelCreateRequest(BaseModel):
68
+ id: str
69
+ name: str
70
+ provider: str
71
+ type: str = "LOCAL"
72
+ tier: str = "BALANCED"
73
+ context_window: int = 32768
74
+ supports_coding: bool = True
75
+ supports_reasoning: bool = True
76
+ supports_vision: bool = False
77
+ supports_tools: bool = True
78
+ quality_score: float = 0.85
79
+ speed_score: float = 0.85
80
+ cost_per_input_token: float = 0.0
81
+ cost_per_output_token: float = 0.0
82
+
83
+
84
+ class RuleCreateRequest(BaseModel):
85
+ id: Optional[str] = None
86
+ name: str
87
+ description: Optional[str] = None
88
+ priority: int = 10
89
+ is_enabled: bool = True
90
+ condition_field: str
91
+ condition_operator: str
92
+ condition_value: str
93
+ action_type: str
94
+ action_target: str
95
+
96
+
97
+ # Helper to convert DB model record to Pydantic ModelMetadata
98
+ def db_model_to_meta(m: ModelRecord) -> ModelMetadata:
99
+ return ModelMetadata(
100
+ id=m.id,
101
+ name=m.name,
102
+ provider=m.provider,
103
+ type=m.type,
104
+ tier=ModelTier(m.tier) if m.tier in [t.value for t in ModelTier] else ModelTier.BALANCED,
105
+ context_window=m.context_window,
106
+ supports_coding=m.supports_coding,
107
+ supports_reasoning=m.supports_reasoning,
108
+ supports_vision=m.supports_vision,
109
+ supports_tools=m.supports_tools,
110
+ quality_score=m.quality_score,
111
+ speed_score=m.speed_score,
112
+ reliability_score=m.reliability_score,
113
+ cost_per_input_token=m.cost_per_input_token,
114
+ cost_per_output_token=m.cost_per_output_token,
115
+ availability=m.availability,
116
+ is_active=m.is_active,
117
+ )
118
+
119
+
120
+ @router.post("/api/route", response_model=Dict[str, Any])
121
+ async def api_route_prompt(req: RouteOnlyRequest, db: AsyncSession = Depends(get_db)):
122
+ """
123
+ Dry-run routing inspection endpoint:
124
+ Analyzes request and produces structured routing decision without invoking provider.
125
+ """
126
+ t_start = time.perf_counter()
127
+ req_id = f"req_{uuid.uuid4().hex[:12]}"
128
+
129
+ # 1. Analyze Request
130
+ analysis = await analyze_request(req.prompt, req.analyzer_mode)
131
+
132
+ # 2. Fetch Models & Policy
133
+ res_m = await db.execute(select(ModelRecord).filter_by(is_active=True))
134
+ db_models = res_m.scalars().all()
135
+ models = [db_model_to_meta(m) for m in db_models]
136
+
137
+ policy_name = req.policy or settings.DEFAULT_ROUTING_POLICY
138
+ res_p = await db.execute(select(RoutingPolicyRecord).filter_by(id=policy_name))
139
+ policy = res_p.scalar_one_or_none()
140
+ weights = {
141
+ "quality_weight": policy.quality_weight if policy else 0.35,
142
+ "cost_weight": policy.cost_weight if policy else 0.25,
143
+ "speed_weight": policy.speed_weight if policy else 0.20,
144
+ "capability_weight": policy.capability_weight if policy else 0.15,
145
+ "reliability_weight": policy.reliability_weight if policy else 0.05,
146
+ }
147
+
148
+ # 3. Fetch Rules & Budget Check
149
+ res_r = await db.execute(select(RoutingRuleRecord).filter_by(is_enabled=True))
150
+ rules = res_r.scalars().all()
151
+ budget_pct, _ = await check_budget_threshold(db, 0.0)
152
+
153
+ # 4. Routing Decision
154
+ decision = route_request(
155
+ analysis=analysis,
156
+ available_models=models,
157
+ policy_weights=weights,
158
+ policy_name=policy_name,
159
+ rules=rules,
160
+ budget_percent=budget_pct,
161
+ request_id=req_id,
162
+ )
163
+
164
+ t_route_ms = (time.perf_counter() - t_start) * 1000.0
165
+
166
+ return {
167
+ "request_id": req_id,
168
+ "analysis": analysis.model_dump(),
169
+ "decision": decision.model_dump(),
170
+ "routing_latency_ms": round(t_route_ms, 2),
171
+ }
172
+
173
+
174
+ @router.post("/api/generate", response_model=Dict[str, Any])
175
+ async def api_generate_response(req: GenerateRequest, db: AsyncSession = Depends(get_db)):
176
+ """
177
+ End-to-End routed execution:
178
+ Analyzes -> Routes -> Executes Selected Model / Fallback -> Records Decision & Metrics.
179
+ """
180
+ t_start = time.perf_counter()
181
+ req_id = f"req_{uuid.uuid4().hex[:12]}"
182
+
183
+ # 1. Analyze Request
184
+ analysis = await analyze_request(req.prompt, req.analyzer_mode)
185
+
186
+ # 2. Fetch Models & Policy
187
+ res_m = await db.execute(select(ModelRecord).filter_by(is_active=True))
188
+ db_models = res_m.scalars().all()
189
+ models = [db_model_to_meta(m) for m in db_models]
190
+
191
+ policy_name = req.policy or settings.DEFAULT_ROUTING_POLICY
192
+ res_p = await db.execute(select(RoutingPolicyRecord).filter_by(id=policy_name))
193
+ policy = res_p.scalar_one_or_none()
194
+ weights = {
195
+ "quality_weight": policy.quality_weight if policy else 0.35,
196
+ "cost_weight": policy.cost_weight if policy else 0.25,
197
+ "speed_weight": policy.speed_weight if policy else 0.20,
198
+ "capability_weight": policy.capability_weight if policy else 0.15,
199
+ "reliability_weight": policy.reliability_weight if policy else 0.05,
200
+ }
201
+
202
+ # 3. Fetch Rules & Budget Check
203
+ res_r = await db.execute(select(RoutingRuleRecord).filter_by(is_enabled=True))
204
+ rules = res_r.scalars().all()
205
+ budget_pct, budget_mode = await check_budget_threshold(db, 0.0)
206
+
207
+ # 4. Route
208
+ decision = route_request(
209
+ analysis=analysis,
210
+ available_models=models,
211
+ policy_weights=weights,
212
+ policy_name=policy_name,
213
+ rules=rules,
214
+ budget_percent=budget_pct,
215
+ request_id=req_id,
216
+ )
217
+
218
+ t_routing_done = time.perf_counter()
219
+ routing_latency_ms = (t_routing_done - t_start) * 1000.0
220
+
221
+ # If force model requested, override
222
+ selected_model_id = req.force_model or decision.selected_model
223
+ target_meta = next((m for m in models if m.id == selected_model_id), None)
224
+ selected_provider_id = target_meta.provider if target_meta else decision.provider
225
+
226
+ # 5. Provider Execution with Fallback & Retries
227
+ provider_resp, fallback_used, orig_model, fallback_reason = await execute_with_fallback(
228
+ prompt=req.prompt,
229
+ selected_model_id=selected_model_id,
230
+ selected_provider_id=selected_provider_id,
231
+ all_models=models,
232
+ system_prompt=req.system_prompt,
233
+ temperature=req.temperature,
234
+ max_retries=req.max_retries,
235
+ )
236
+
237
+ total_latency_ms = (time.perf_counter() - t_start) * 1000.0
238
+
239
+ # 6. Calculate actual costs
240
+ actual_model_meta = next((m for m in models if m.id == provider_resp.model), target_meta)
241
+ cost_in_rate = actual_model_meta.cost_per_input_token if actual_model_meta else 0.0
242
+ cost_out_rate = actual_model_meta.cost_per_output_token if actual_model_meta else 0.0
243
+ actual_cost = round((provider_resp.input_tokens * cost_in_rate) + (provider_resp.output_tokens * cost_out_rate), 6)
244
+
245
+ # Baseline cost calculation
246
+ res_base = await db.execute(select(ModelRecord).filter_by(id=settings.BASELINE_MODEL_ID))
247
+ base_m = res_base.scalar_one_or_none()
248
+ base_cost = 0.0
249
+ if base_m:
250
+ base_cost = round((provider_resp.input_tokens * base_m.cost_per_input_token) + (provider_resp.output_tokens * base_m.cost_per_output_token), 6)
251
+ saved_cost = max(0.0, round(base_cost - actual_cost, 6))
252
+
253
+ # Update budget spend
254
+ await check_budget_threshold(db, actual_cost)
255
+
256
+ # 7. Persistence
257
+ req_record = RequestRecord(
258
+ request_id=req_id,
259
+ prompt=req.prompt,
260
+ task_type=analysis.task_type.value,
261
+ complexity=analysis.complexity,
262
+ context_size=analysis.context_size,
263
+ reasoning_required=analysis.reasoning_required,
264
+ coding_required=analysis.coding_required,
265
+ routing_policy=policy_name,
266
+ selected_model=provider_resp.model or selected_model_id,
267
+ provider=provider_resp.provider or selected_provider_id,
268
+ status="FALLBACK" if fallback_used else ("FAILED" if provider_resp.error else "SUCCESS"),
269
+ fallback_used=fallback_used,
270
+ original_model=orig_model,
271
+ fallback_reason=fallback_reason,
272
+ input_tokens=provider_resp.input_tokens,
273
+ output_tokens=provider_resp.output_tokens,
274
+ total_tokens=provider_resp.total_tokens,
275
+ estimated_cost=actual_cost,
276
+ baseline_cost=base_cost,
277
+ cost_saved=saved_cost,
278
+ routing_latency_ms=round(routing_latency_ms, 2),
279
+ provider_latency_ms=round(provider_resp.provider_latency_ms, 2),
280
+ total_latency_ms=round(total_latency_ms, 2),
281
+ time_to_first_token_ms=provider_resp.time_to_first_token_ms,
282
+ )
283
+ db.add(req_record)
284
+
285
+ decision_record = RoutingDecisionRecord(
286
+ decision_id=decision.decision_id,
287
+ request_id=req_id,
288
+ selected_model=decision.selected_model,
289
+ confidence=decision.confidence,
290
+ reasons=decision.reasons,
291
+ candidate_scores=[s.model_dump() for s in decision.candidate_scores],
292
+ rejected_candidates=decision.rejected_candidates,
293
+ policy_used=policy_name,
294
+ )
295
+ db.add(decision_record)
296
+
297
+ resp_record = ResponseRecord(
298
+ response_id=f"resp_{uuid.uuid4().hex[:12]}",
299
+ request_id=req_id,
300
+ model_id=provider_resp.model or selected_model_id,
301
+ provider=provider_resp.provider or selected_provider_id,
302
+ content=provider_resp.content,
303
+ finish_reason=provider_resp.finish_reason,
304
+ is_mock=provider_resp.is_mock,
305
+ )
306
+ db.add(resp_record)
307
+ await db.commit()
308
+
309
+ # 8. Structured Observability Event
310
+ log_router_event(
311
+ event_name="request_completed",
312
+ request_id=req_id,
313
+ model=provider_resp.model,
314
+ provider=provider_resp.provider,
315
+ duration_ms=round(total_latency_ms, 2),
316
+ metadata={"tokens": provider_resp.total_tokens, "cost": actual_cost, "fallback": fallback_used},
317
+ )
318
+
319
+ return {
320
+ "request_id": req_id,
321
+ "analysis": analysis.model_dump(),
322
+ "decision": decision.model_dump(),
323
+ "response": provider_resp.model_dump(),
324
+ "metrics": {
325
+ "routing_latency_ms": round(routing_latency_ms, 2),
326
+ "provider_latency_ms": round(provider_resp.provider_latency_ms, 2),
327
+ "total_latency_ms": round(total_latency_ms, 2),
328
+ "estimated_cost_usd": actual_cost,
329
+ "baseline_cost_usd": base_cost,
330
+ "cost_saved_usd": saved_cost,
331
+ "fallback_used": fallback_used,
332
+ "fallback_reason": fallback_reason,
333
+ },
334
+ }
335
+
336
+
337
+ # Models & Registry Endpoints
338
+ @router.get("/api/models")
339
+ async def list_models(db: AsyncSession = Depends(get_db)):
340
+ res = await db.execute(select(ModelRecord).order_by(ModelRecord.tier, ModelRecord.name))
341
+ return res.scalars().all()
342
+
343
+
344
+ @router.post("/api/models")
345
+ async def create_model(req: ModelCreateRequest, db: AsyncSession = Depends(get_db)):
346
+ record = ModelRecord(
347
+ id=req.id,
348
+ name=req.name,
349
+ provider=req.provider,
350
+ type=req.type,
351
+ tier=req.tier,
352
+ context_window=req.context_window,
353
+ supports_coding=req.supports_coding,
354
+ supports_reasoning=req.supports_reasoning,
355
+ supports_vision=req.supports_vision,
356
+ supports_tools=req.supports_tools,
357
+ quality_score=req.quality_score,
358
+ speed_score=req.speed_score,
359
+ cost_per_input_token=req.cost_per_input_token,
360
+ cost_per_output_token=req.cost_per_output_token,
361
+ )
362
+ db.add(record)
363
+ await db.commit()
364
+ return record
365
+
366
+
367
+ # Providers Endpoint (Never returning secrets)
368
+ @router.get("/api/providers")
369
+ async def list_providers(db: AsyncSession = Depends(get_db)):
370
+ providers_info = []
371
+ for pid, p in provider_registry.list_providers().items():
372
+ health = await p.check_health()
373
+ providers_info.append({
374
+ "id": pid,
375
+ "name": p.name,
376
+ "status": health.get("status", "READY"),
377
+ "base_url": p.base_url,
378
+ "message": health.get("message"),
379
+ "models_available": health.get("models_available", []),
380
+ "credentials_status": "Loaded from environment (.env)" if getattr(p, "api_key", None) or pid in ["mock", "ollama"] else "Not configured",
381
+ })
382
+ return providers_info
383
+
384
+
385
+ # Live Traffic & Decisions
386
+ @router.get("/api/traffic")
387
+ async def get_traffic_feed(limit: int = 50, db: AsyncSession = Depends(get_db)):
388
+ res = await db.execute(select(RequestRecord).order_by(desc(RequestRecord.timestamp)).limit(limit))
389
+ return res.scalars().all()
390
+
391
+
392
+ @router.get("/api/traffic/export")
393
+ async def export_traffic(
394
+ format: str = Query("json", pattern="^(csv|json)$"),
395
+ db: AsyncSession = Depends(get_db),
396
+ ):
397
+ res = await db.execute(select(RequestRecord).order_by(desc(RequestRecord.timestamp)))
398
+ records = res.scalars().all()
399
+ fields = [
400
+ "timestamp",
401
+ "request_id",
402
+ "prompt_preview",
403
+ "task_type",
404
+ "complexity",
405
+ "selected_model",
406
+ "input_tokens",
407
+ "output_tokens",
408
+ "total_tokens",
409
+ "cost_saved",
410
+ "total_latency_ms",
411
+ ]
412
+ rows = [
413
+ {
414
+ "timestamp": record.timestamp.isoformat() if record.timestamp else "",
415
+ "request_id": record.request_id,
416
+ "prompt_preview": record.prompt[:200],
417
+ "task_type": record.task_type,
418
+ "complexity": record.complexity,
419
+ "selected_model": record.selected_model,
420
+ "input_tokens": record.input_tokens,
421
+ "output_tokens": record.output_tokens,
422
+ "total_tokens": record.total_tokens,
423
+ "cost_saved": record.cost_saved,
424
+ "total_latency_ms": record.total_latency_ms,
425
+ }
426
+ for record in records
427
+ ]
428
+
429
+ if format == "json":
430
+ return JSONResponse(
431
+ content=rows,
432
+ headers={"Content-Disposition": "attachment; filename=traffic-export.json"},
433
+ )
434
+
435
+ output = io.StringIO()
436
+ writer = csv.DictWriter(output, fieldnames=fields)
437
+ writer.writeheader()
438
+ writer.writerows(rows)
439
+ return StreamingResponse(
440
+ iter([output.getvalue()]),
441
+ media_type="text/csv",
442
+ headers={"Content-Disposition": "attachment; filename=traffic-export.csv"},
443
+ )
444
+
445
+
446
+ @router.get("/api/decisions/{decision_id_or_request_id}")
447
+ async def get_decision_details(decision_id_or_request_id: str, db: AsyncSession = Depends(get_db)):
448
+ res_dec = await db.execute(
449
+ select(RoutingDecisionRecord).filter(
450
+ (RoutingDecisionRecord.decision_id == decision_id_or_request_id)
451
+ | (RoutingDecisionRecord.request_id == decision_id_or_request_id)
452
+ )
453
+ )
454
+ dec = res_dec.scalar_one_or_none()
455
+ if not dec:
456
+ raise HTTPException(status_code=404, detail="Decision not found")
457
+
458
+ res_req = await db.execute(select(RequestRecord).filter_by(request_id=dec.request_id))
459
+ req = res_req.scalar_one_or_none()
460
+
461
+ res_resp = await db.execute(select(ResponseRecord).filter_by(request_id=dec.request_id))
462
+ resp = res_resp.scalar_one_or_none()
463
+
464
+ return {
465
+ "decision": dec,
466
+ "request": req,
467
+ "response": resp,
468
+ }
469
+
470
+
471
+ # Feedback
472
+ @router.post("/api/feedback")
473
+ async def submit_feedback(req: FeedbackSubmitRequest, db: AsyncSession = Depends(get_db)):
474
+ res = await db.execute(select(RequestRecord).filter_by(request_id=req.request_id))
475
+ request_rec = res.scalar_one_or_none()
476
+ if not request_rec:
477
+ raise HTTPException(status_code=404, detail="Request ID not found")
478
+
479
+ feedback = FeedbackRecord(
480
+ id=f"fb_{uuid.uuid4().hex[:10]}",
481
+ request_id=req.request_id,
482
+ model_id=request_rec.selected_model,
483
+ task_type=request_rec.task_type,
484
+ rating=1 if req.rating > 0 else -1,
485
+ comment=req.comment,
486
+ )
487
+ db.add(feedback)
488
+ await db.commit()
489
+ return {"status": "success", "feedback_id": feedback.id}
490
+
491
+
492
+ # Analytics Endpoints
493
+ @router.get("/api/analytics")
494
+ async def get_analytics(db: AsyncSession = Depends(get_db)):
495
+ return await get_system_analytics(db)
496
+
497
+
498
+ @router.get("/api/analytics/savings")
499
+ async def get_savings(baseline_model: Optional[str] = None, db: AsyncSession = Depends(get_db)):
500
+ return await calculate_cost_savings(db, baseline_model)
501
+
502
+
503
+ # Rules Endpoints
504
+ @router.get("/api/rules")
505
+ async def list_rules(db: AsyncSession = Depends(get_db)):
506
+ res = await db.execute(select(RoutingRuleRecord).order_by(desc(RoutingRuleRecord.priority)))
507
+ return res.scalars().all()
508
+
509
+
510
+ @router.post("/api/rules")
511
+ async def create_rule(req: RuleCreateRequest, db: AsyncSession = Depends(get_db)):
512
+ rule_id = req.id or f"rule-{uuid.uuid4().hex[:8]}"
513
+ record = RoutingRuleRecord(
514
+ id=rule_id,
515
+ name=req.name,
516
+ description=req.description,
517
+ priority=req.priority,
518
+ is_enabled=req.is_enabled,
519
+ condition_field=req.condition_field,
520
+ condition_operator=req.condition_operator,
521
+ condition_value=req.condition_value,
522
+ action_type=req.action_type,
523
+ action_target=req.action_target,
524
+ )
525
+ db.add(record)
526
+ await db.commit()
527
+ return record
528
+
529
+
530
+ @router.put("/api/rules/{rule_id}")
531
+ async def update_rule(rule_id: str, req: RuleCreateRequest, db: AsyncSession = Depends(get_db)):
532
+ res = await db.execute(select(RoutingRuleRecord).filter_by(id=rule_id))
533
+ rule = res.scalar_one_or_none()
534
+ if not rule:
535
+ raise HTTPException(status_code=404, detail="Rule not found")
536
+ rule.name = req.name
537
+ rule.description = req.description
538
+ rule.priority = req.priority
539
+ rule.is_enabled = req.is_enabled
540
+ rule.condition_field = req.condition_field
541
+ rule.condition_operator = req.condition_operator
542
+ rule.condition_value = req.condition_value
543
+ rule.action_type = req.action_type
544
+ rule.action_target = req.action_target
545
+ await db.commit()
546
+ return rule
547
+
548
+
549
+ @router.delete("/api/rules/{rule_id}")
550
+ async def delete_rule(rule_id: str, db: AsyncSession = Depends(get_db)):
551
+ res = await db.execute(select(RoutingRuleRecord).filter_by(id=rule_id))
552
+ rule = res.scalar_one_or_none()
553
+ if not rule:
554
+ raise HTTPException(status_code=404, detail="Rule not found")
555
+ await db.delete(rule)
556
+ await db.commit()
557
+ return {"status": "deleted", "id": rule_id}
558
+
559
+
560
+ # Policies
561
+ @router.get("/api/policies")
562
+ async def list_policies(db: AsyncSession = Depends(get_db)):
563
+ res = await db.execute(select(RoutingPolicyRecord))
564
+ return res.scalars().all()
565
+
566
+
567
+ # Budgets
568
+ @router.get("/api/budgets")
569
+ async def get_budget(db: AsyncSession = Depends(get_db)):
570
+ res = await db.execute(select(BudgetRecord).filter_by(id="default"))
571
+ return res.scalar_one_or_none()
572
+
573
+
574
+ # Experiments
575
+ @router.get("/api/experiments")
576
+ async def list_experiments(db: AsyncSession = Depends(get_db)):
577
+ res = await db.execute(select(ExperimentRecord))
578
+ return res.scalars().all()
579
+
580
+
581
+ # Health Check
582
+ @router.get("/api/health")
583
+ async def get_health():
584
+ return {
585
+ "status": "healthy",
586
+ "service": "Model Router AI Control Room",
587
+ "version": "1.0.0",
588
+ "environment": settings.APP_ENV,
589
+ }
app/budgets/manager.py ADDED
@@ -0,0 +1,39 @@
1
+ from typing import Tuple
2
+ from sqlalchemy import select
3
+ from sqlalchemy.ext.asyncio import AsyncSession
4
+ from app.storage.models import BudgetRecord
5
+
6
+
7
+ async def check_budget_threshold(
8
+ session: AsyncSession,
9
+ cost_to_add: float = 0.0,
10
+ ) -> Tuple[float, str]:
11
+ """
12
+ Checks current budget utilization percentage and determines intervention mode:
13
+ - NORMAL (< 80%)
14
+ - OPTIMIZE_80 (>= 80% and < 95%)
15
+ - CHEAP_95 (>= 95% and < 100%)
16
+ - BLOCK_100 (>= 100%)
17
+ """
18
+ res = await session.execute(select(BudgetRecord).filter_by(id="default"))
19
+ budget = res.scalar_one_or_none()
20
+ if not budget:
21
+ return 0.0, "NORMAL"
22
+
23
+ monthly_limit = budget.monthly_limit or 100.0
24
+ current_spend = budget.current_monthly_spend + cost_to_add
25
+ percent = (current_spend / monthly_limit) * 100.0 if monthly_limit > 0 else 0.0
26
+
27
+ mode = "NORMAL"
28
+ if percent >= 100.0:
29
+ mode = "BLOCK_100"
30
+ elif percent >= 95.0:
31
+ mode = "CHEAP_95"
32
+ elif percent >= 80.0:
33
+ mode = "OPTIMIZE_80"
34
+
35
+ budget.current_monthly_spend = current_spend
36
+ budget.intervention_mode = mode
37
+ await session.commit()
38
+
39
+ return round(percent, 2), mode