mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-22 12:22:20 +00:00
Compare commits
4 Commits
fix/refund
...
embeddings
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5b0829090e | ||
|
|
4ed86b0091 | ||
|
|
b327c0dee1 | ||
|
|
470c32aebf |
@@ -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
routstr/model_mappings.json
Normal file
7
routstr/model_mappings.json
Normal file
@@ -0,0 +1,7 @@
|
||||
{
|
||||
"manual_model_mappings": {
|
||||
"mappings": {
|
||||
"text-embedding-ada-002-v2": "text-embedding-ada-002"
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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.lower())
|
||||
"""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]:
|
||||
|
||||
@@ -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
ui/lib/api/services/modelMappings.ts
Normal file
78
ui/lib/api/services/modelMappings.ts
Normal 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',
|
||||
{}
|
||||
);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user