Compare commits

...
8 changed files with 630 additions and 66 deletions
+121 -1
View File
@@ -5,7 +5,7 @@ from pathlib import Path
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from fastapi.responses import HTMLResponse, RedirectResponse
from pydantic import BaseModel
from pydantic import BaseModel, Field
from sqlmodel import select
from ..payment.models import _row_to_model, list_models
@@ -3165,3 +3165,123 @@ async def get_log_dates_api(request: Request) -> dict[str, object]:
continue
return {"dates": dates}
class ModelMappingRequest(BaseModel):
from_model: str = Field(..., alias="from")
to: str
class ModelMappingUpdateRequest(BaseModel):
to: str
@admin_router.get("/api/model-mappings", dependencies=[Depends(require_admin_api)])
async def get_model_mappings(request: Request) -> dict[str, str]:
from ..proxy import _manual_model_mappings
return _manual_model_mappings
@admin_router.post("/api/model-mappings", dependencies=[Depends(require_admin_api)])
async def create_model_mapping(request: Request, mapping: ModelMappingRequest) -> dict[str, str]:
import json
import os
from ..proxy import _manual_model_mappings, load_manual_model_mappings
mappings_file = os.path.join(os.path.dirname(os.path.dirname(__file__)), "model_mappings.json")
try:
if os.path.exists(mappings_file):
with open(mappings_file, "r") as f:
data = json.load(f)
else:
data = {"manual_model_mappings": {"mappings": {}}}
data["manual_model_mappings"]["mappings"][mapping.from_model.lower()] = mapping.to.lower()
with open(mappings_file, "w") as f:
json.dump(data, f, indent=2)
load_manual_model_mappings()
return _manual_model_mappings
except Exception as e:
raise HTTPException(status_code=500, detail=f"Failed to create mapping: {str(e)}")
@admin_router.put("/api/model-mappings/{from_model}", dependencies=[Depends(require_admin_api)])
async def update_model_mapping(request: Request, from_model: str, mapping: ModelMappingUpdateRequest) -> dict[str, str]:
import json
import os
from ..proxy import _manual_model_mappings, load_manual_model_mappings
mappings_file = os.path.join(os.path.dirname(os.path.dirname(__file__)), "model_mappings.json")
try:
if os.path.exists(mappings_file):
with open(mappings_file, "r") as f:
data = json.load(f)
else:
data = {"manual_model_mappings": {"mappings": {}}}
if from_model.lower() not in data["manual_model_mappings"]["mappings"]:
raise HTTPException(status_code=404, detail="Mapping not found")
data["manual_model_mappings"]["mappings"][from_model.lower()] = mapping.to.lower()
with open(mappings_file, "w") as f:
json.dump(data, f, indent=2)
load_manual_model_mappings()
return _manual_model_mappings
except HTTPException:
raise
except Exception as e:
raise HTTPException(status_code=500, detail=f"Failed to update mapping: {str(e)}")
@admin_router.delete("/api/model-mappings/{from_model}", dependencies=[Depends(require_admin_api)])
async def delete_model_mapping(request: Request, from_model: str) -> dict[str, str]:
import json
import os
from ..proxy import _manual_model_mappings, load_manual_model_mappings
mappings_file = os.path.join(os.path.dirname(os.path.dirname(__file__)), "model_mappings.json")
try:
if os.path.exists(mappings_file):
with open(mappings_file, "r") as f:
data = json.load(f)
else:
data = {"manual_model_mappings": {"mappings": {}}}
if from_model.lower() not in data["manual_model_mappings"]["mappings"]:
raise HTTPException(status_code=404, detail="Mapping not found")
del data["manual_model_mappings"]["mappings"][from_model.lower()]
with open(mappings_file, "w") as f:
json.dump(data, f, indent=2)
load_manual_model_mappings()
return _manual_model_mappings
except HTTPException:
raise
except Exception as e:
raise HTTPException(status_code=500, detail=f"Failed to delete mapping: {str(e)}")
@admin_router.post("/api/model-mappings/reload", dependencies=[Depends(require_admin_api)])
async def reload_model_mappings(request: Request) -> dict[str, object]:
from ..proxy import _manual_model_mappings, load_manual_model_mappings
try:
load_manual_model_mappings()
return {"ok": True, "mappings": _manual_model_mappings}
except Exception as e:
raise HTTPException(status_code=500, detail=f"Failed to reload mappings: {str(e)}")
+7
View File
@@ -0,0 +1,7 @@
{
"manual_model_mappings": {
"mappings": {
"text-embedding-ada-002-v2": "text-embedding-ada-002"
}
}
}
+62 -2
View File
@@ -5,6 +5,7 @@ from pydantic.v1 import BaseModel
from ..core import get_logger
from ..core.db import AsyncSession
from ..core.settings import settings
from .price import sats_usd_price
logger = get_logger(__name__)
@@ -64,6 +65,56 @@ async def calculate_cost( # todo: can be sync
)
return cost_data
usage_data = response_data["usage"]
usd_cost = 0.0
# Prioritize cost_details.upstream_inference_cost
if "cost_details" in usage_data:
usd_cost = float(
usage_data["cost_details"].get("upstream_inference_cost", 0) or 0
)
# Fallback to cost field if upstream_inference_cost is 0
if usd_cost == 0 and "cost" in usage_data:
try:
usd_cost = float(usage_data.get("cost", 0) or 0)
except Exception:
pass
if usd_cost > 0:
try:
sats_per_usd = 1.0 / sats_usd_price()
cost_in_sats = usd_cost * sats_per_usd
cost_in_msats = math.ceil(cost_in_sats * 1000)
logger.info(
"Using cost from usage data/details",
extra={
"usd_cost": usd_cost,
"cost_in_sats": cost_in_sats,
"cost_in_msats": cost_in_msats,
"model": response_data.get("model", "unknown"),
},
)
return CostData(
base_msats=-1,
input_msats=-1, # Cost field doesn't break down by token type
output_msats=-1,
total_msats=cost_in_msats,
)
except Exception as e:
logger.warning(
"Error calculating cost from usage data",
extra={
"error": str(e),
"usd_cost": usd_cost,
"model": response_data.get("model", "unknown"),
},
)
# Fall through to token-based calculation
MSATS_PER_1K_INPUT_TOKENS: float = (
float(settings.fixed_per_1k_input_tokens) * 1000.0
)
@@ -129,10 +180,19 @@ async def calculate_cost( # todo: can be sync
)
return cost_data
input_tokens = response_data.get("usage", {}).get("prompt_tokens", 0)
output_tokens = response_data.get("usage", {}).get("completion_tokens", 0)
input_tokens = usage_data.get("prompt_tokens", 0)
output_tokens = usage_data.get("completion_tokens", 0)
# added for response api
input_tokens = (
input_tokens if input_tokens != 0 else usage_data.get("input_tokens", 0)
)
output_tokens = (
output_tokens if output_tokens != 0 else usage_data.get("output_tokens", 0)
)
input_msats = round(input_tokens / 1000 * MSATS_PER_1K_INPUT_TOKENS, 3)
output_msats = round(output_tokens / 1000 * MSATS_PER_1K_OUTPUT_TOKENS, 3)
token_based_cost = math.ceil(input_msats + output_msats)
+26 -6
View File
@@ -93,12 +93,32 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis
try:
async with httpx.AsyncClient() as client:
response = await client.get(f"{base_url}/models", timeout=30)
response.raise_for_status()
data = response.json()
models_response, embeddings_response = await asyncio.gather(
client.get(f"{base_url}/models", timeout=30),
client.get(f"{base_url}/embeddings/models", timeout=30),
return_exceptions=True,
)
def process_models_response(
response: httpx.Response | BaseException,
) -> list[dict]:
if not isinstance(response, BaseException):
response.raise_for_status()
data = response.json()
return [
model
for model in data.get("data", [])
if ":free" not in model.get("id", "").lower()
]
return []
models_data: list[dict] = []
for model in data.get("data", []):
models_data.extend(process_models_response(models_response))
models_data.extend(process_models_response(embeddings_response))
# Apply source filter and exclusions
filtered_models = []
for model in models_data:
model_id = model.get("id", "")
if source_filter:
@@ -116,9 +136,9 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis
if not _has_valid_pricing(model):
continue
models_data.append(model)
filtered_models.append(model)
return models_data
return filtered_models
except Exception as e:
logger.error(f"Error (async) fetching models from OpenRouter API: {e}")
return []
+43 -4
View File
@@ -1,4 +1,5 @@
import json
import os
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Request
@@ -33,6 +34,23 @@ _upstreams: list[BaseUpstreamProvider] = []
_model_instances: dict[str, Model] = {} # All aliases -> Model
_provider_map: dict[str, BaseUpstreamProvider] = {} # All aliases -> Provider
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
_manual_model_mappings: dict[str, str] = {} # Manual model_id mappings loaded from JSON
def load_manual_model_mappings() -> None:
"""Load manual model mappings from JSON file."""
global _manual_model_mappings
try:
mappings_file = os.path.join(os.path.dirname(__file__), "model_mappings.json")
if os.path.exists(mappings_file):
with open(mappings_file, "r") as f:
data = json.load(f)
_manual_model_mappings = data.get("manual_model_mappings", {}).get("mappings", {})
else:
_manual_model_mappings = {}
except Exception as e:
logger.error(f"Failed to load manual model mappings: {e}")
_manual_model_mappings = {}
async def initialize_upstreams() -> None:
@@ -40,6 +58,7 @@ async def initialize_upstreams() -> None:
global _upstreams
_upstreams = await init_upstreams()
logger.info(f"Initialized {len(_upstreams)} upstream providers")
load_manual_model_mappings()
await refresh_model_maps()
@@ -51,6 +70,7 @@ async def reinitialize_upstreams() -> None:
"Re-initialized upstream providers from admin action",
extra={"provider_count": len(_upstreams)},
)
load_manual_model_mappings()
await refresh_model_maps()
@@ -64,13 +84,32 @@ def get_upstreams() -> list[BaseUpstreamProvider]:
def get_model_instance(model_id: str) -> Model | None:
"""Get Model instance by ID from global cache."""
return _model_instances.get(model_id)
"""Get Model instance by ID from global cache, with manual mapping fallback."""
model = _model_instances.get(model_id)
if model is not None:
return model
mapped_model_id = _manual_model_mappings.get(model_id.lower())
if mapped_model_id:
return _model_instances.get(mapped_model_id.lower())
return None
def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None:
"""Get UpstreamProvider for model ID from global cache."""
return _provider_map.get(model_id)
"""Get UpstreamProvider for model ID from global cache, with manual mapping fallback."""
# First try direct lookup
provider = _provider_map.get(model_id)
if provider is not None:
return provider
# Try manual mapping as fallback
mapped_model_id = _manual_model_mappings.get(model_id)
if mapped_model_id:
logger.debug(f"Using manual mapping for provider: {model_id} -> {mapped_model_id}")
return _provider_map.get(mapped_model_id)
return None
def get_unique_models() -> list[Model]:
+71 -51
View File
@@ -734,51 +734,54 @@ class BaseUpstreamProvider:
await client.aclose()
return mapped_error
if path.endswith("chat/completions"):
client_wants_streaming = False
if request_body:
try:
request_data = json.loads(request_body)
client_wants_streaming = request_data.get("stream", False)
logger.debug(
"Chat completion request analysis",
extra={
"client_wants_streaming": client_wants_streaming,
"model": request_data.get("model", "unknown"),
"key_hash": key.hashed_key[:8] + "...",
},
)
except json.JSONDecodeError:
logger.warning(
"Failed to parse request body JSON for streaming detection"
)
# Handle endpoints that require cost calculation and payment adjustment
if path.endswith("chat/completions") or path.endswith("embeddings"):
if path.endswith("chat/completions"):
client_wants_streaming = False
if request_body:
try:
request_data = json.loads(request_body)
client_wants_streaming = request_data.get("stream", False)
logger.debug(
"Chat completion request analysis",
extra={
"client_wants_streaming": client_wants_streaming,
"model": request_data.get("model", "unknown"),
"key_hash": key.hashed_key[:8] + "...",
},
)
except json.JSONDecodeError:
logger.warning(
"Failed to parse request body JSON for streaming detection"
)
content_type = response.headers.get("content-type", "")
upstream_is_streaming = "text/event-stream" in content_type
is_streaming = client_wants_streaming and upstream_is_streaming
content_type = response.headers.get("content-type", "")
upstream_is_streaming = "text/event-stream" in content_type
is_streaming = client_wants_streaming and upstream_is_streaming
logger.debug(
"Response type analysis",
extra={
"is_streaming": is_streaming,
"client_wants_streaming": client_wants_streaming,
"upstream_is_streaming": upstream_is_streaming,
"content_type": content_type,
"key_hash": key.hashed_key[:8] + "...",
},
)
if is_streaming and response.status_code == 200:
result = await self.handle_streaming_chat_completion(
response, key, max_cost_for_model
logger.debug(
"Response type analysis",
extra={
"is_streaming": is_streaming,
"client_wants_streaming": client_wants_streaming,
"upstream_is_streaming": upstream_is_streaming,
"content_type": content_type,
"key_hash": key.hashed_key[:8] + "...",
},
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
background_tasks.add_task(client.aclose)
result.background = background_tasks
return result
elif response.status_code == 200:
if is_streaming and response.status_code == 200:
result = await self.handle_streaming_chat_completion(
response, key, max_cost_for_model
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
background_tasks.add_task(client.aclose)
result.background = background_tasks
return result
# Handle both non-streaming chat completions and embeddings
if response.status_code == 200:
try:
return await self.handle_non_streaming_chat_completion(
response, key, session, max_cost_for_model
@@ -1519,9 +1522,9 @@ class BaseUpstreamProvider:
error_response.headers["X-Cashu"] = refund_token
return error_response
if path.endswith("chat/completions"):
if path.endswith("chat/completions") or path.endswith("embeddings"):
logger.debug(
"Processing chat completion response",
"Processing completion/embeddings response",
extra={"path": path, "amount": amount, "unit": unit},
)
@@ -1770,15 +1773,32 @@ class BaseUpstreamProvider:
async def _fetch_openrouter_models(self) -> list[dict]:
"""Fetch models from OpenRouter API."""
url = "https://openrouter.ai/api/v1/models"
embeddings_url = "https://openrouter.ai/api/v1/embeddings/models"
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.get(url)
response.raise_for_status()
models = response.json()
return [
model
for model in models.get("data", [])
if ":free" not in model.get("id", "").lower()
]
models_response, embeddings_response = await asyncio.gather(
client.get(url), client.get(embeddings_url), return_exceptions=True
)
all_models = []
def process_models_response(
response: httpx.Response | BaseException,
) -> list[dict]:
if not isinstance(response, BaseException):
response.raise_for_status()
data = response.json()
return [
model
for model in data.get("data", [])
if ":free" not in model.get("id", "").lower()
]
return []
all_models.extend(process_models_response(models_response))
all_models.extend(process_models_response(embeddings_response))
return all_models
async def _fetch_provider_models(self) -> dict:
"""Fetch models from provider's API."""
+222 -2
View File
@@ -10,16 +10,26 @@ import { SiteHeader } from '@/components/site-header';
import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs';
import { useQuery } from '@tanstack/react-query';
import { AdminService } from '@/lib/api/services/admin';
import { ModelMappingService } from '@/lib/api/services/modelMappings';
import { Skeleton } from '@/components/ui/skeleton';
import { AlertCircle, Users, Globe } from 'lucide-react';
import { Alert, AlertDescription } from '@/components/ui/alert';
import { Badge } from '@/components/ui/badge';
import { useMemo, useState } from 'react';
import React, { useMemo, useState } from 'react';
import type { Model } from '@/lib/api/schemas/models';
import { groupAndSortModelsByProvider } from '@/lib/utils/modelSort';
import { Button } from '@/components/ui/button';
import { Input } from '@/components/ui/input';
import { Card, CardContent, CardHeader, CardTitle } from '@/components/ui/card';
import { Trash2, Plus, Edit2, Save, X } from 'lucide-react';
export default function ModelsPage() {
const [filteredModels, setFilteredModels] = useState<Model[]>([]);
const [modelMappings, setModelMappings] = useState<Record<string, string>>(
{}
);
const [editingMapping, setEditingMapping] = useState<string | null>(null);
const [newMapping, setNewMapping] = useState({ from: '', to: '' });
const {
data: modelsData,
@@ -31,6 +41,23 @@ export default function ModelsPage() {
refetchOnWindowFocus: false,
});
const {
data: mappingsData,
isLoading: isLoadingMappings,
error: mappingsError,
refetch: refetchMappings,
} = useQuery({
queryKey: ['model-mappings'],
queryFn: () => ModelMappingService.getModelMappings(),
refetchOnWindowFocus: false,
});
React.useEffect(() => {
if (mappingsData) {
setModelMappings(mappingsData);
}
}, [mappingsData]);
const { models = [], groups = [] } = modelsData || {};
const groupedModels = useMemo(() => {
@@ -67,6 +94,40 @@ export default function ModelsPage() {
});
}, [groupedModels, groupDataMap, groups]);
const handleAddMapping = async () => {
if (!newMapping.from || !newMapping.to) return;
try {
await ModelMappingService.createModelMapping({
from: newMapping.from,
to: newMapping.to,
});
setNewMapping({ from: '', to: '' });
refetchMappings();
} catch (error) {
console.error('Failed to add mapping:', error);
}
};
const handleDeleteMapping = async (from: string) => {
try {
await ModelMappingService.deleteModelMapping(from);
refetchMappings();
} catch (error) {
console.error('Failed to delete mapping:', error);
}
};
const handleUpdateMapping = async (from: string, to: string) => {
try {
await ModelMappingService.updateModelMapping(from, { to });
setEditingMapping(null);
refetchMappings();
} catch (error) {
console.error('Failed to update mapping:', error);
}
};
return (
<SidebarProvider>
<AppSidebar variant='inset' />
@@ -81,8 +142,9 @@ export default function ModelsPage() {
</div>
<Tabs defaultValue='manage' className='w-full'>
<TabsList className='grid w-full grid-cols-3'>
<TabsList className='grid w-full grid-cols-4'>
<TabsTrigger value='manage'>Manage Models</TabsTrigger>
<TabsTrigger value='mappings'>Model Mappings</TabsTrigger>
{/*<TabsTrigger value='test-basic'>Basic Testing</TabsTrigger>
<TabsTrigger value='test-api'>API Endpoints</TabsTrigger> */}
</TabsList>
@@ -267,6 +329,164 @@ export default function ModelsPage() {
)}
</TabsContent>
<TabsContent value='mappings' className='space-y-4'>
<div className='text-muted-foreground text-sm'>
Manage model ID mappings to redirect requests from one model
to another. This is useful for maintaining compatibility with
legacy model names or creating aliases.
</div>
{isLoadingMappings ? (
<div className='space-y-4'>
<Skeleton className='h-[200px] w-full' />
</div>
) : mappingsError ? (
<Alert variant='destructive'>
<AlertCircle className='h-4 w-4' />
<AlertDescription>
Failed to load model mappings. Please try refreshing the
page.
</AlertDescription>
</Alert>
) : (
<div className='space-y-6'>
<Card>
<CardHeader>
<CardTitle className='flex items-center gap-2'>
<Plus className='h-5 w-5' />
Add New Model Mapping
</CardTitle>
</CardHeader>
<CardContent>
<div className='grid grid-cols-1 gap-4 md:grid-cols-3'>
<Input
placeholder='From model ID'
value={newMapping.from}
onChange={(e) =>
setNewMapping({
...newMapping,
from: e.target.value,
})
}
/>
<Input
placeholder='To model ID'
value={newMapping.to}
onChange={(e) =>
setNewMapping({
...newMapping,
to: e.target.value,
})
}
/>
<Button
onClick={handleAddMapping}
disabled={!newMapping.from || !newMapping.to}
className='w-full'
>
<Plus className='mr-2 h-4 w-4' />
Add Mapping
</Button>
</div>
</CardContent>
</Card>
<Card>
<CardHeader>
<CardTitle>Current Model Mappings</CardTitle>
</CardHeader>
<CardContent>
{Object.keys(modelMappings).length === 0 ? (
<div className='text-muted-foreground py-8 text-center'>
No model mappings configured
</div>
) : (
<div className='space-y-3'>
{Object.entries(modelMappings).map(([from, to]) => (
<div
key={from}
className='flex items-center justify-between gap-4 rounded-lg border p-4'
>
<div className='grid flex-1 grid-cols-1 gap-4 md:grid-cols-2'>
<div>
<label className='text-muted-foreground text-sm font-medium'>
From
</label>
<div className='font-mono text-sm'>
{from}
</div>
</div>
<div>
<label className='text-muted-foreground text-sm font-medium'>
To
</label>
{editingMapping === from ? (
<div className='flex items-center gap-2'>
<Input
defaultValue={to}
id={`edit-${from}`}
className='text-sm'
/>
<Button
size='sm'
onClick={() => {
const input =
document.getElementById(
`edit-${from}`
) as HTMLInputElement;
handleUpdateMapping(
from,
input.value
);
}}
>
<Save className='h-4 w-4' />
</Button>
<Button
size='sm'
variant='outline'
onClick={() =>
setEditingMapping(null)
}
>
<X className='h-4 w-4' />
</Button>
</div>
) : (
<div className='font-mono text-sm'>
{to}
</div>
)}
</div>
</div>
{editingMapping !== from && (
<div className='flex items-center gap-2'>
<Button
size='sm'
variant='outline'
onClick={() => setEditingMapping(from)}
>
<Edit2 className='h-4 w-4' />
</Button>
<Button
size='sm'
variant='destructive'
onClick={() => handleDeleteMapping(from)}
>
<Trash2 className='h-4 w-4' />
</Button>
</div>
)}
</div>
))}
</div>
)}
</CardContent>
</Card>
</div>
)}
</TabsContent>
<TabsContent value='test-basic' className='space-y-4'>
<div className='text-muted-foreground text-sm'>
Test model credentials and connectivity with basic chat
+78
View File
@@ -0,0 +1,78 @@
import { apiClient } from '../client';
import { z } from 'zod';
export const ModelMappingSchema = z.object({
from: z.string(),
to: z.string(),
});
export const CreateModelMappingSchema = z.object({
from: z.string(),
to: z.string(),
});
export const UpdateModelMappingSchema = z.object({
to: z.string(),
});
export const ModelMappingsResponseSchema = z.record(z.string());
export const ReloadMappingsResponseSchema = z.object({
ok: z.boolean(),
mappings: z.record(z.string()),
});
export type ModelMapping = z.infer<typeof ModelMappingSchema>;
export type CreateModelMapping = z.infer<typeof CreateModelMappingSchema>;
export type UpdateModelMapping = z.infer<typeof UpdateModelMappingSchema>;
export type ModelMappingsResponse = z.infer<typeof ModelMappingsResponseSchema>;
export type ReloadMappingsResponse = z.infer<
typeof ReloadMappingsResponseSchema
>;
export class ModelMappingService {
static async getModelMappings(): Promise<ModelMappingsResponse> {
return await apiClient.get<ModelMappingsResponse>(
'/admin/api/model-mappings'
);
}
static async createModelMapping(
data: CreateModelMapping
): Promise<ModelMappingsResponse> {
return await apiClient.post<ModelMappingsResponse>(
'/admin/api/model-mappings',
{
from: data.from,
to: data.to,
}
);
}
static async updateModelMapping(
fromModel: string,
data: UpdateModelMapping
): Promise<ModelMappingsResponse> {
return await apiClient.put<ModelMappingsResponse>(
`/admin/api/model-mappings/${encodeURIComponent(fromModel)}`,
{
to: data.to,
}
);
}
static async deleteModelMapping(
fromModel: string
): Promise<ModelMappingsResponse> {
return await apiClient.delete<ModelMappingsResponse>(
`/admin/api/model-mappings/${encodeURIComponent(fromModel)}`
);
}
static async reloadModelMappings(): Promise<ReloadMappingsResponse> {
return await apiClient.post<ReloadMappingsResponse>(
'/admin/api/model-mappings/reload',
{}
);
}
}