Compare commits

..
Author SHA1 Message Date
9qeklajc 5b0829090e fmt 2025-12-22 22:01:28 +01:00
9qeklajc 4ed86b0091 update model mapping to the ui 2025-12-22 21:40:06 +01:00
9qeklajc b327c0dee1 clean up 2025-12-22 21:12:27 +01:00
9qeklajc 470c32aebf add static model mapping 2025-12-22 21:10:45 +01:00
34 changed files with 808 additions and 1740 deletions
+37
View File
@@ -0,0 +1,37 @@
import os
import openai
client = openai.OpenAI(
api_key=os.environ["CASHU_TOKEN"],
base_url=os.environ.get("ROUTSTR_API_URL", "https://api.routstr.com/v1"),
# base_url="http://roustrjfsdgfiueghsklchg.onion/v1",
# client=httpx.AsyncClient(
# proxies={"http": "socks5://localhost:9050"},
# ), # to use onion proxy (tor)
)
history: list = []
def chat() -> None:
while True:
user_msg = {"role": "user", "content": input("\nYou: ")}
history.append(user_msg)
ai_msg = {"role": "assistant", "content": ""}
for chunk in client.chat.completions.create(
model=os.environ.get("MODEL", "openai/gpt-4o-mini"),
messages=history,
stream=True,
):
if len(chunk.choices) > 0:
content = chunk.choices[0].delta.content
if content is not None:
ai_msg["content"] += content
print(content, end="", flush=True)
print()
history.append(ai_msg)
if __name__ == "__main__":
chat()
-11
View File
@@ -1,11 +0,0 @@
import os
import httpx
# Use your Cashu token or API key as the Bearer token,
# cashu token is hashed on the server and acts as an Temporary API key
headers = {"Authorization": f"Bearer {os.environ.get('TOKEN')}"}
base_url = os.environ.get("API_URL", "https://api.routstr.com/v1")
resp = httpx.get(f"{base_url}/balance/info", headers=headers)
print(resp.json())
-15
View File
@@ -1,15 +0,0 @@
import os
import httpx
# Send a Cashu token to the /create endpoint to get a persistent API key
token = os.environ.get("TOKEN")
if not token:
print("Please set TOKEN environment variable with a Cashu token")
exit(1)
base_url = os.environ.get("API_URL", "https://api.routstr.com/v1")
resp = httpx.get(f"{base_url}/balance/create", params={"initial_balance_token": token})
print(resp.json())
-12
View File
@@ -1,12 +0,0 @@
import os
import httpx
# Use your Cashu token or API key as the Bearer token
headers = {"Authorization": f"Bearer {os.environ.get('TOKEN')}"}
base_url = os.environ.get("API_URL", "https://api.routstr.com/v1")
resp = httpx.post(f"{base_url}/balance/refund", headers=headers)
print("Refund successful!")
print(resp.json())
-16
View File
@@ -1,16 +0,0 @@
import os
import httpx
# Use your Cashu token or API key as the Bearer token
headers = {"Authorization": f"Bearer {os.environ.get('TOKEN')}"}
base_url = os.environ.get("API_URL", "https://api.routstr.com/v1")
# The Cashu token to top up with
cashu_token = input("Enter Cashu token to top up: ")
resp = httpx.post(
f"{base_url}/balance/topup", headers=headers, json={"cashu_token": cashu_token}
)
print(resp.json())
-15
View File
@@ -1,15 +0,0 @@
import os
from openai import OpenAI
client = OpenAI(
api_key=os.environ.get("TOKEN"),
base_url=os.environ.get("API_URL", "https://api.routstr.com/v1"),
)
response = client.chat.completions.create(
model=os.environ.get("MODEL", "gpt-5-nano"),
messages=[{"role": "user", "content": "Hello!"}],
)
print(response.choices[0].message.content)
-19
View File
@@ -1,19 +0,0 @@
import os
import httpx
from openai import OpenAI
client = OpenAI(
api_key=os.environ.get("TOKEN", ""),
base_url=os.environ.get("API_URL", "https://api.routstr.com/v1"),
)
for model in client.models.list():
print(model.id)
# OR
models = httpx.get(
f"{client.base_url}/v1/models",
headers={"Authorization": f"Bearer {client.api_key}"},
).json()
-31
View File
@@ -1,31 +0,0 @@
import os
from openai import OpenAI
client = OpenAI(
api_key=os.environ.get("TOKEN"),
base_url=os.environ.get("API_URL", "https://api.routstr.com/v1"),
)
conversation = [] # type: ignore
# First turn
response1 = client.responses.create( # type: ignore
model="o4-mini",
input="Hi, my name is Alice.",
conversation=conversation,
)
print("Response 1:", response1.output)
# Note: The 'conversation' parameter might need to be constructed differently
# depending on exact SDK/API spec. Typically, you pass back the previous turn's data.
# Assuming the SDK manages or returns a conversation object/ID:
# conversation.append(response1)
# Second turn - demonstrating intent, actual implementation depends on strict API spec
# response2 = client.responses.create(
# model="openai/gpt-4o-mini",
# input="What is my name?",
# conversation=conversation,
# )
# print("Response 2:", response2.output)
-17
View File
@@ -1,17 +0,0 @@
import os
from openai import OpenAI
# The OpenAI SDK handles the 'responses' endpoint if it's updated to the latest version
# and the base_url points to a compatible proxy like Routstr.
client = OpenAI(
api_key=os.environ.get("TOKEN"),
base_url=os.environ.get("API_URL", "https://api.routstr.com/v1"),
)
response = client.responses.create(
model="gpt-5-mini",
input="Tell me a three sentence bedtime story about a unicorn.",
)
print(response.output)
-20
View File
@@ -1,20 +0,0 @@
import os
from openai import OpenAI
client = OpenAI(
api_key=os.environ.get("TOKEN"),
base_url=os.environ.get("API_URL", "https://api.routstr.com/v1"),
)
stream = client.responses.create(
model="claude-4.5-sonnet",
input="Write a short poem about rust.",
stream=True,
)
for event in stream:
# Note: Depending on the SDK version and response structure,
# you might access event.output_delta or similar fields
print(event, end="", flush=True)
print()
-16
View File
@@ -1,16 +0,0 @@
import os
from openai import OpenAI
client = OpenAI(
api_key=os.environ.get("TOKEN"),
base_url=os.environ.get("API_URL", "https://api.routstr.com/v1"),
)
response = client.responses.create(
model="gpt-5-mini",
input="What is the latest news about AI?",
tools=[{"type": "web_search"}], # type: ignore
)
print(response.output)
-28
View File
@@ -1,28 +0,0 @@
import os
from openai import OpenAI
client = OpenAI(
api_key=os.environ.get("TOKEN"),
base_url=os.environ.get("API_URL", "https://api.routstr.com/v1"),
)
messages = []
while True:
messages.append({"role": "user", "content": input("\nYou: ")})
stream = client.chat.completions.create(
model=os.environ.get("MODEL", "gpt-5.1-mini"),
messages=messages, # type: ignore
stream=True,
)
print("AI: ", end="")
response_content = ""
for chunk in stream:
if content := chunk.choices[0].delta.content: # type: ignore
print(content, end="", flush=True)
response_content += content
print()
messages.append({"role": "assistant", "content": response_content})
-20
View File
@@ -1,20 +0,0 @@
import os
import httpx
from openai import OpenAI
# Requires `pip install "httpx[socks]"` and a running Tor proxy on port 9050
client = OpenAI(
api_key=os.environ.get("TOKEN"),
base_url=os.environ.get("ONION_URL", "http://roustrjfsdgfiueghsklchg.onion/v1"),
http_client=httpx.Client(proxies="socks5://localhost:9050"),
)
print(
client.chat.completions.create(
model="openai/gpt-4o-mini",
messages=[{"role": "user", "content": "Hello from Tor!"}],
)
.choices[0]
.message.content
)
@@ -1,37 +0,0 @@
"""alias-ids
Revision ID: b9667ffc5701
Revises: lightning_invoices
Create Date: 2025-12-25 19:30:44.673350
"""
import sqlalchemy as sa
import sqlmodel
from alembic import op
# revision identifiers, used by Alembic.
revision = "b9667ffc5701"
down_revision = "lightning_invoices"
branch_labels = None
depends_on = None
def upgrade() -> None:
# ### commands auto generated by Alembic ###
op.add_column(
"models",
sa.Column("canonical_slug", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
)
op.add_column(
"models",
sa.Column("alias_ids", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
)
# ### end Alembic commands ###
def downgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
op.drop_column("models", "alias_ids")
op.drop_column("models", "canonical_slug")
# ### end Alembic commands ###
-1
View File
@@ -73,7 +73,6 @@ packages = ["routstr"]
[tool.ruff.lint]
select = ["E", "F", "I"]
ignore = ["E501"]
exclude = ["examples"]
[tool.mypy]
python_version = "3.11"
+6 -7
View File
@@ -206,20 +206,19 @@ def create_model_mappings(
alias: str, model: "Model", provider: "BaseUpstreamProvider"
) -> None:
"""Set alias to model/provider if not set or if new model is preferred."""
alias_lower = alias.lower()
existing_model = model_instances.get(alias_lower)
existing_model = model_instances.get(alias)
if not existing_model:
# No existing mapping, set it
model_instances[alias_lower] = model
provider_map[alias_lower] = provider
model_instances[alias] = model
provider_map[alias] = provider
else:
# Check if candidate should replace existing
existing_provider = provider_map[alias_lower]
existing_provider = provider_map[alias]
if should_prefer_model(
model, provider, existing_model, existing_provider, alias
):
model_instances[alias_lower] = model
provider_map[alias_lower] = provider
model_instances[alias] = model
provider_map[alias] = provider
def process_provider_models(
upstream: "BaseUpstreamProvider", is_openrouter: bool = False
+231 -79
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
@@ -1493,8 +1493,20 @@ class ModelCreate(BaseModel):
per_request_limits: dict[str, object] | None = None
top_provider: dict[str, object] | None = None
upstream_provider_id: int | None = None
canonical_slug: str | None = None
alias_ids: list[str] | None = None
enabled: bool = True
class ModelUpdate(BaseModel):
id: str
name: str
description: str
created: int
context_length: int
architecture: dict[str, object]
pricing: dict[str, object]
per_request_limits: dict[str, object] | None = None
top_provider: dict[str, object] | None = None
upstream_provider_id: int | None = None
enabled: bool = True
@@ -2404,80 +2416,44 @@ async def admin_upstream_providers(request: Request) -> str:
"/api/upstream-providers/{provider_id}/models",
dependencies=[Depends(require_admin_api)],
)
async def upsert_provider_model(
async def create_provider_model(
provider_id: int, payload: ModelCreate
) -> dict[str, object]:
print(payload)
logger.info(
f"UPSERT_PROVIDER_MODEL called: provider_id={provider_id}, model_id={payload.id}"
)
async with create_session() as session:
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
# Try to get existing model
existing_row = await session.get(ModelRow, (payload.id, provider_id))
exists = await session.get(ModelRow, (payload.id, provider_id))
if exists:
raise HTTPException(
status_code=409,
detail="Model with this ID already exists for this provider",
)
if existing_row:
# Update existing model
logger.info(f"Updating existing model: {payload.id}")
existing_row.name = payload.name
existing_row.description = payload.description
existing_row.created = int(payload.created)
existing_row.context_length = int(payload.context_length)
existing_row.architecture = json.dumps(payload.architecture)
existing_row.pricing = json.dumps(payload.pricing)
existing_row.sats_pricing = None
existing_row.per_request_limits = (
row = ModelRow(
id=payload.id,
name=payload.name,
description=payload.description,
created=int(payload.created),
context_length=int(payload.context_length),
architecture=json.dumps(payload.architecture),
pricing=json.dumps(payload.pricing),
sats_pricing=None,
per_request_limits=(
json.dumps(payload.per_request_limits)
if payload.per_request_limits is not None
else None
)
existing_row.top_provider = (
),
top_provider=(
json.dumps(payload.top_provider) if payload.top_provider else None
)
existing_row.canonical_slug = payload.canonical_slug
existing_row.alias_ids = (
json.dumps(payload.alias_ids) if payload.alias_ids else None
)
existing_row.enabled = payload.enabled
session.add(existing_row)
await session.commit()
await session.refresh(existing_row)
row = existing_row
else:
# Create new model
logger.info(f"Creating new model: {payload.id}")
row = ModelRow(
id=payload.id,
name=payload.name,
description=payload.description,
created=int(payload.created),
context_length=int(payload.context_length),
architecture=json.dumps(payload.architecture),
pricing=json.dumps(payload.pricing),
sats_pricing=None,
per_request_limits=(
json.dumps(payload.per_request_limits)
if payload.per_request_limits is not None
else None
),
top_provider=(
json.dumps(payload.top_provider) if payload.top_provider else None
),
canonical_slug=payload.canonical_slug,
alias_ids=(
json.dumps(payload.alias_ids) if payload.alias_ids else None
),
upstream_provider_id=provider_id,
enabled=payload.enabled,
)
session.add(row)
await session.commit()
await session.refresh(row)
),
upstream_provider_id=provider_id,
enabled=payload.enabled,
)
session.add(row)
await session.commit()
await session.refresh(row)
await refresh_model_maps()
return _row_to_model(
@@ -2485,20 +2461,6 @@ async def upsert_provider_model(
).dict() # type: ignore
@admin_router.patch(
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
dependencies=[Depends(require_admin_api)],
)
async def update_provider_model_legacy(
provider_id: int, model_id: str, payload: ModelCreate
) -> dict[str, object]:
"""Legacy PATCH endpoint - redirects to upsert POST endpoint for backward compatibility."""
logger.info(
f"LEGACY_PATCH_UPDATE called: provider_id={provider_id}, model_id={model_id}"
)
return await upsert_provider_model(provider_id, payload)
@admin_router.get(
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
dependencies=[Depends(require_admin_api)],
@@ -2519,6 +2481,76 @@ async def get_provider_model(provider_id: int, model_id: str) -> dict[str, objec
).dict() # type: ignore
@admin_router.patch(
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
dependencies=[Depends(require_admin_api)],
)
async def update_provider_model(
provider_id: int, model_id: str, payload: ModelUpdate
) -> dict[str, object]:
if payload.id != model_id:
raise HTTPException(status_code=400, detail="Path id does not match payload id")
async with create_session() as session:
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
row = await session.get(ModelRow, (model_id, provider_id))
if not row:
raise HTTPException(
status_code=404, detail="Model not found for this provider"
)
row.name = payload.name
row.description = payload.description
row.created = int(payload.created)
row.context_length = int(payload.context_length)
row.architecture = json.dumps(payload.architecture)
row.pricing = json.dumps(payload.pricing)
row.sats_pricing = None
row.per_request_limits = (
json.dumps(payload.per_request_limits)
if payload.per_request_limits is not None
else None
)
row.top_provider = (
json.dumps(payload.top_provider) if payload.top_provider else None
)
was_disabled = not row.enabled
row.enabled = payload.enabled
session.add(row)
await session.commit()
await session.refresh(row)
if was_disabled and payload.enabled:
from ..payment.models import _cleanup_enabled_models_once
try:
await _cleanup_enabled_models_once()
except Exception as e:
logger.warning(
f"Failed to run model cleanup after enabling: {e}",
extra={"model_id": model_id, "error": str(e)},
)
await refresh_model_maps()
return _row_to_model(
row, apply_provider_fee=True, provider_fee=provider.provider_fee
).dict() # type: ignore
@admin_router.put(
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
dependencies=[Depends(require_admin_api)],
)
async def update_provider_model_put(
provider_id: int, model_id: str, payload: ModelUpdate
) -> dict[str, object]:
return await update_provider_model(provider_id, model_id, payload)
@admin_router.delete(
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
dependencies=[Depends(require_admin_api)],
@@ -3133,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)}")
-4
View File
@@ -68,10 +68,6 @@ class ModelRow(SQLModel, table=True): # type: ignore
sats_pricing: str | None = Field(default=None)
per_request_limits: str | None = Field(default=None)
top_provider: str | None = Field(default=None)
canonical_slug: str | None = Field(default=None, description="Canonical model slug")
alias_ids: str | None = Field(
default=None, description="JSON array of model alias IDs"
)
enabled: bool = Field(default=True, description="Whether this model is enabled")
upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models")
+12 -10
View File
@@ -14,6 +14,7 @@ from ..balance import balance_router, deprecated_wallet_router
from ..discovery import providers_cache_refresher, providers_router
from ..nip91 import announce_provider
from ..payment.models import (
cleanup_enabled_models_periodically,
models_router,
update_sats_pricing,
)
@@ -48,6 +49,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
nip91_task = None
providers_task = None
models_refresh_task = None
models_cleanup_task = None
model_maps_refresh_task = None
try:
@@ -62,11 +64,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
async with create_session() as session:
s = await SettingsService.initialize(session)
if not s.admin_password:
logger.warning(
f"Admin password is not set. Visit {s.http_url or 'http://localhost:8000'}/admin to set the password."
)
# Apply app metadata from settings
try:
app.title = s.name
@@ -83,17 +80,13 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
_update_prices_task = asyncio.create_task(_update_prices())
_initialize_upstreams_task = asyncio.create_task(initialize_upstreams())
# ensure both setup tasks complete
await asyncio.gather(
_update_prices_task, _initialize_upstreams_task, return_exceptions=True
)
btc_price_task = asyncio.create_task(update_prices_periodically())
pricing_task = asyncio.create_task(update_sats_pricing())
if global_settings.models_refresh_interval_seconds > 0:
models_refresh_task = asyncio.create_task(
refresh_upstreams_models_periodically(get_upstreams())
)
models_cleanup_task = asyncio.create_task(cleanup_enabled_models_periodically())
model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically())
payout_task = asyncio.create_task(periodic_payout())
if global_settings.nsec:
@@ -101,6 +94,11 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
if global_settings.providers_refresh_interval_seconds > 0:
providers_task = asyncio.create_task(providers_cache_refresher())
# ensure both setup tasks complete
await asyncio.gather(
_update_prices_task, _initialize_upstreams_task, return_exceptions=True
)
yield
except asyncio.CancelledError:
@@ -127,6 +125,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
providers_task.cancel()
if models_refresh_task is not None:
models_refresh_task.cancel()
if models_cleanup_task is not None:
models_cleanup_task.cancel()
if model_maps_refresh_task is not None:
model_maps_refresh_task.cancel()
@@ -144,6 +144,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
tasks_to_wait.append(providers_task)
if models_refresh_task is not None:
tasks_to_wait.append(models_refresh_task)
if models_cleanup_task is not None:
tasks_to_wait.append(models_cleanup_task)
if model_maps_refresh_task is not None:
tasks_to_wait.append(model_maps_refresh_task)
+1 -10
View File
@@ -55,16 +55,7 @@ class LoggingMiddleware(BaseHTTPMiddleware):
"headers": {
k: v
for k, v in request.headers.items()
if k.lower()
not in [
"authorization",
"x-cashu",
"cookie",
"cf-connecting-ip",
"cf-ipcountry",
"x-forwarded-for",
"x-real-ip",
]
if k.lower() not in ["authorization", "x-cashu", "cookie"]
},
"body_size": len(request_body) if request_body else 0,
},
+7
View File
@@ -0,0 +1,7 @@
{
"manual_model_mappings": {
"mappings": {
"text-embedding-ada-002-v2": "text-embedding-ada-002"
}
}
}
-4
View File
@@ -191,10 +191,6 @@ async def calculate_cost( # todo: can be sync
output_tokens if output_tokens != 0 else usage_data.get("output_tokens", 0)
)
# added for response api
input_tokens = input_tokens if input_tokens != 0 else response_data.get("usage", {}).get("input_tokens", 0)
output_tokens = output_tokens if output_tokens != 0 else response_data.get("usage", {}).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)
+112 -11
View File
@@ -5,9 +5,10 @@ import random
import httpx
from fastapi import APIRouter, Depends
from pydantic.v1 import BaseModel
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from ..core.db import ModelRow, get_session
from ..core.db import ModelRow, create_session, get_session
from ..core.logging import get_logger
from ..core.settings import settings
from .price import sats_usd_price
@@ -184,7 +185,6 @@ def _row_to_model(
enabled=row.enabled,
upstream_provider_id=row.upstream_provider_id,
canonical_slug=getattr(row, "canonical_slug", None),
alias_ids=json.loads(row.alias_ids) if row.alias_ids else None,
)
if apply_provider_fee:
@@ -382,7 +382,7 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
async def _update_sats_pricing_once() -> None:
"""Update sats pricing once for all provider models (in-memory only)."""
from ..proxy import get_upstreams, refresh_model_maps
from ..proxy import get_upstreams
upstreams = get_upstreams()
sats_to_usd = sats_usd_price()
@@ -399,7 +399,6 @@ async def _update_sats_pricing_once() -> None:
if updated_count > 0:
logger.info("Updated sats pricing", extra={"models_updated": updated_count})
await refresh_model_maps()
async def update_sats_pricing() -> None:
@@ -410,13 +409,7 @@ async def update_sats_pricing() -> None:
except Exception:
pass
try:
await _update_sats_pricing_once()
except Exception as e:
logger.warning(
"Initial sats pricing update failed (will retry in loop)",
extra={"error": str(e)},
)
await _update_sats_pricing_once()
while True:
try:
@@ -440,6 +433,114 @@ async def update_sats_pricing() -> None:
logger.error(f"Error updating sats pricing: {e}")
async def cleanup_enabled_models_periodically() -> None:
"""Background task to clean up enabled models that match upstream pricing.
When model is enabled (enabled=True), remove it from DB if it matches upstream pricing.
Keep it in DB only if pricing differs from upstream or if it's disabled.
"""
interval = getattr(
settings, "models_cleanup_interval_seconds", 300
) # 5 minutes default
if not interval or interval <= 0:
return
while True:
try:
await _cleanup_enabled_models_once()
except asyncio.CancelledError:
break
except Exception as e:
logger.error(
"Error during enabled models cleanup",
extra={"error": str(e), "error_type": type(e).__name__},
)
try:
jitter = max(0.0, float(interval) * 0.1)
await asyncio.sleep(interval + random.uniform(0, jitter))
except asyncio.CancelledError:
break
async def _cleanup_enabled_models_once() -> None:
"""Clean up enabled models that match upstream pricing."""
from ..proxy import get_upstreams
async with create_session() as session:
# Get all enabled models from DB
result = await session.exec(
select(ModelRow).where(
ModelRow.enabled, # Only enabled models
)
)
db_models = result.all()
if not db_models:
return
upstreams = get_upstreams()
models_to_remove = []
for db_model in db_models:
# Find corresponding upstream model
upstream_model = None
for upstream in upstreams:
upstream_model = upstream.get_cached_model_by_id(db_model.id)
if upstream_model:
break
if not upstream_model:
continue
# Compare pricing to see if they match
db_pricing = json.loads(db_model.pricing)
upstream_pricing = upstream_model.pricing.dict()
# Check if pricing matches (with small tolerance for float comparison)
pricing_matches = _pricing_matches(db_pricing, upstream_pricing)
if pricing_matches:
models_to_remove.append(db_model)
logger.info(
f"Removing enabled model {db_model.id} - matches upstream pricing",
extra={"model_id": db_model.id},
)
# Remove models that match upstream pricing
for model in models_to_remove:
await session.delete(model)
if models_to_remove:
await session.commit()
logger.info(
f"Cleaned up {len(models_to_remove)} enabled models that match upstream pricing"
)
def _pricing_matches(
db_pricing: dict, upstream_pricing: dict, tolerance: float = 0.0
) -> bool:
"""Check if pricing dictionaries match within tolerance."""
keys_to_compare = [
"prompt",
"completion",
"request",
"image",
"web_search",
"internal_reasoning",
]
for key in keys_to_compare:
db_val = int(float(db_pricing.get(key, 0.0)) * 1000000)
upstream_val = int(float(upstream_pricing.get(key, 0.0)) * 1000000)
if abs(db_val - upstream_val) > tolerance:
return False
return True
@models_router.get("/v1/models")
@models_router.get("/models", include_in_schema=False)
async def models(session: AsyncSession = Depends(get_session)) -> dict:
+92 -74
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.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.lower())
"""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]:
@@ -80,27 +119,31 @@ def get_unique_models() -> list[Model]:
async def refresh_model_maps() -> None:
"""Refresh global model and provider maps using the cost-based algorithm."""
from sqlalchemy.orm import selectinload
global _model_instances, _provider_map, _unique_models
# Gather database overrides and disabled models
async with create_session() as session:
# Fetch all providers with their models in a single logical operation
query = select(UpstreamProviderRow).options(
selectinload(UpstreamProviderRow.models) # type: ignore
result = await session.exec(select(ModelRow).where(ModelRow.enabled))
override_rows = result.all()
provider_result = await session.exec(select(UpstreamProviderRow))
providers_by_id = {p.id: p for p in provider_result.all()}
overrides_by_id: dict[str, tuple[ModelRow, float]] = {
row.id: (
row,
providers_by_id[row.upstream_provider_id].provider_fee
if row.upstream_provider_id in providers_by_id
else 1.01,
)
for row in override_rows
if row.upstream_provider_id is not None
}
disabled_result = await session.exec(
select(ModelRow.id).where(ModelRow.enabled == False) # noqa: E712
)
result = await session.exec(query)
provider_rows = result.all()
overrides_by_id: dict[str, tuple[ModelRow, float]] = {}
disabled_model_ids: set[str] = set()
for provider in provider_rows:
for model in provider.models:
if model.enabled:
overrides_by_id[model.id] = (model, provider.provider_fee)
else:
disabled_model_ids.add(model.id)
disabled_model_ids = {row for row in disabled_result.all()}
_model_instances, _provider_map, _unique_models = create_model_mappings(
upstreams=_upstreams,
@@ -137,14 +180,20 @@ async def proxy(
"unauthorized", "Unauthorized", 401, request=request
)
is_responses_api = path.startswith("v1/responses") or path.startswith("responses")
logger.info( # TODO: move to middleware, async
"Received proxy request",
extra={
"method": request.method,
"path": path,
"client_host": request.client.host if request.client else "unknown",
"user_agent": request.headers.get("user-agent", "unknown")[:100],
},
)
request_body = await request.body()
request_body_dict = parse_request_body_json(request_body, path)
if is_responses_api:
model_id = extract_model_from_responses_request(request_body_dict)
else:
model_id = request_body_dict.get("model", "unknown")
model_id = request_body_dict.get("model", "unknown")
model_obj = get_model_instance(model_id)
if not model_obj:
@@ -170,14 +219,9 @@ async def proxy(
check_token_balance(headers, request_body_dict, max_cost_for_model)
if x_cashu := headers.get("x-cashu", None):
if is_responses_api:
return await upstream.handle_x_cashu_responses(
request, x_cashu, path, max_cost_for_model, model_obj
)
else:
return await upstream.handle_x_cashu(
request, x_cashu, path, max_cost_for_model, model_obj
)
return await upstream.handle_x_cashu(
request, x_cashu, path, max_cost_for_model, model_obj
)
elif auth := headers.get("authorization", None):
key = await get_bearer_token_key(headers, path, session, auth)
@@ -192,36 +236,28 @@ async def proxy(
)
logger.debug("Processing unauthenticated GET request", extra={"path": path})
# TODO: why is this needed? can we remove it?
headers = upstream.prepare_headers(dict(request.headers))
return await upstream.forward_get_request(request, path, headers)
# Only pay for request if we have request body data (for completions endpoints)
if request_body_dict:
await pay_for_request(key, max_cost_for_model, session)
# Prepare headers for upstream
headers = upstream.prepare_headers(dict(request.headers))
if is_responses_api:
response = await upstream.forward_responses_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
else:
response = await upstream.forward_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
# Forward to upstream and handle response
response = await upstream.forward_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
if response.status_code != 200:
await revert_pay_for_request(key, session, max_cost_for_model)
@@ -324,24 +360,6 @@ async def get_bearer_token_key(
raise
def extract_model_from_responses_request(request_body_dict: dict[str, Any]) -> str:
if model := request_body_dict.get("model"):
return model
if input_data := request_body_dict.get("input"):
if isinstance(input_data, dict) and (model := input_data.get("model")):
return model
if request_body_dict.get("messages"):
return "unknown"
logger.warning(
"No model found in Responses API request",
extra={"body_keys": list(request_body_dict.keys())}
)
return "unknown"
def parse_request_body_json(request_body: bytes, path: str) -> dict[str, Any]:
request_body_dict = {}
if request_body:
+1 -1142
View File
File diff suppressed because it is too large Load Diff
+6 -23
View File
@@ -1,6 +1,5 @@
from __future__ import annotations
import asyncio
import os
import re
from typing import TYPE_CHECKING
@@ -8,7 +7,6 @@ from typing import TYPE_CHECKING
if TYPE_CHECKING:
from ..core.settings import Settings
from sqlmodel import select
from ..core import get_logger
from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session
@@ -147,17 +145,6 @@ async def refresh_upstreams_models_periodically(
f"Error refreshing models for {upstream.base_url}",
extra={"error": str(e), "error_type": type(e).__name__},
)
try:
from ..payment.models import _update_sats_pricing_once
await _update_sats_pricing_once()
except Exception as e:
logger.warning(f"Failed to update pricing after model refresh: {e}")
from ..proxy import refresh_model_maps
await refresh_model_maps()
except asyncio.CancelledError:
break
except Exception as e:
@@ -179,6 +166,8 @@ async def init_upstreams() -> list[BaseUpstreamProvider]:
Seeds database with providers from settings if empty, then loads and instantiates
provider instances from database records, and refreshes their models cache.
"""
from sqlmodel import select
from ..core.settings import settings
async with create_session() as session:
@@ -194,16 +183,16 @@ async def init_upstreams() -> list[BaseUpstreamProvider]:
result = await session.exec(select(UpstreamProviderRow))
existing_providers = result.all()
async def _init_single_provider(
provider_row: UpstreamProviderRow,
) -> BaseUpstreamProvider | None:
upstreams: list[BaseUpstreamProvider] = []
for provider_row in existing_providers:
if not provider_row.enabled:
logger.debug(f"Skipping disabled provider: {provider_row.base_url}")
return None
continue
provider = _instantiate_provider(provider_row)
if provider:
await provider.refresh_models_cache()
upstreams.append(provider)
logger.debug(
f"Initialized {provider_row.provider_type} provider",
extra={
@@ -211,12 +200,6 @@ async def init_upstreams() -> list[BaseUpstreamProvider]:
"models_cached": len(provider.get_cached_models()),
},
)
return provider
return None
tasks = [_init_single_provider(row) for row in existing_providers]
results = await asyncio.gather(*tasks)
upstreams = [p for p in results if p is not None]
return upstreams
+1 -7
View File
@@ -50,13 +50,7 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider):
async def fetch_models(self) -> list[Model]:
"""Fetch all OpenRouter models."""
models_data = await async_fetch_openrouter_models()
models = [Model(**model) for model in models_data] # type: ignore
# manual alias for openai/text-embedding-ada-002 due to openrouter api bug
for model in models:
if model.id == "openai/text-embedding-ada-002":
model.alias_ids = ["text-embedding-ada-002-v2"]
break
return models
return [Model(**model) for model in models_data] # type: ignore
async def get_balance(self) -> float | None:
"""Get the current account balance from OpenRouter.
-97
View File
@@ -1,97 +0,0 @@
import json
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from httpx import AsyncClient
@pytest.mark.integration
@pytest.mark.asyncio
async def test_proxy_embeddings_endpoint(authenticated_client: AsyncClient) -> None:
"""Test the embeddings endpoint proxy functionality"""
test_payload = {
"model": "text-embedding-ada-002",
"input": "The quick brown fox",
}
mock_response_data = {
"object": "list",
"data": [
{"object": "embedding", "embedding": [0.0023, -0.0012, 0.0045], "index": 0}
],
"model": "text-embedding-ada-002",
"usage": {"prompt_tokens": 5, "total_tokens": 5},
}
with patch("httpx.AsyncClient.send") as mock_send:
# Create a proper async generator for iter_bytes
async def mock_iter_bytes(*args: Any, **kwargs: Any) -> Any:
yield json.dumps(mock_response_data).encode()
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
# Use MagicMock for synchronous .json() method
mock_response.json = MagicMock(return_value=mock_response_data)
mock_response.iter_bytes = mock_iter_bytes
mock_response.aiter_bytes = mock_iter_bytes
mock_send.return_value = mock_response
# Make POST request to embeddings endpoint
response = await authenticated_client.post("/v1/embeddings", json=test_payload)
assert response.status_code == 200
response_data = response.json()
assert response_data["object"] == "list"
assert len(response_data["data"]) == 1
assert response_data["data"][0]["object"] == "embedding"
# Verify request was forwarded
mock_send.assert_called_once()
forwarded_request = mock_send.call_args[0][0]
# Verify the path ends with embeddings
# Note: forwarded path might be full URL
assert str(forwarded_request.url).endswith("embeddings")
@pytest.mark.integration
@pytest.mark.asyncio
async def test_model_case_insensitivity(authenticated_client: AsyncClient) -> None:
"""Test that model lookups are case insensitive"""
# We'll use a mixed-case model ID that should match the lowercase one in the system
# We assume 'gpt-3.5-turbo' is available in the mock env/database
test_payload = {
"model": "GPT-3.5-TURBO",
"messages": [{"role": "user", "content": "Hello"}],
}
with patch("httpx.AsyncClient.send") as mock_send:
mock_response_data = {
"id": "chatcmpl-123",
"object": "chat.completion",
"choices": [{"message": {"content": "Hi"}}],
"usage": {"total_tokens": 10},
}
async def mock_iter_bytes(*args: Any, **kwargs: Any) -> Any:
yield json.dumps(mock_response_data).encode()
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_response.json = MagicMock(return_value=mock_response_data)
mock_response.iter_bytes = mock_iter_bytes
mock_response.aiter_bytes = mock_iter_bytes
mock_send.return_value = mock_response
response = await authenticated_client.post(
"/v1/chat/completions", json=test_payload
)
assert response.status_code == 200
+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
-4
View File
@@ -41,10 +41,6 @@ export function ModelSearchFilter({
model.name.toLowerCase().includes(query) ||
model.full_name.toLowerCase().includes(query) ||
model.provider.toLowerCase().includes(query) ||
(model.alias_ids &&
model.alias_ids.some((alias) =>
alias.toLowerCase().includes(query)
)) ||
(model.description &&
model.description.toLowerCase().includes(query)) ||
model.modelType.toLowerCase().includes(query)
-2
View File
@@ -542,7 +542,6 @@ export function ModelSelector({
top_provider: null,
upstream_provider_id: providerId,
enabled: model.isEnabled,
alias_ids: model.alias_ids || null,
};
setModelDialogState({
@@ -586,7 +585,6 @@ export function ModelSelector({
top_provider: null,
upstream_provider_id: providerId,
enabled: model.isEnabled,
alias_ids: model.alias_ids || null,
};
setModelDialogState({
-1
View File
@@ -31,7 +31,6 @@ export const ModelSchema = z.object({
// API key type indicators
has_own_api_key: z.boolean(),
api_key_type: z.string(), // "individual" or "group"
alias_ids: z.array(z.string()).nullable().optional(),
});
// Schema for a model with additional provider-specific settings
+2 -5
View File
@@ -121,7 +121,6 @@ export interface AdminModelAsModel {
soft_deleted?: boolean;
has_own_api_key: boolean;
api_key_type: string;
alias_ids?: string[] | null;
}
export interface AdminModelGroup {
@@ -208,7 +207,6 @@ export class AdminService {
soft_deleted: !adminModel.enabled,
has_own_api_key: false,
api_key_type: 'group',
alias_ids: adminModel.alias_ids,
};
}
@@ -356,9 +354,8 @@ export class AdminService {
original: data.pricing,
converted: payload.pricing,
});
// Use the same POST endpoint for both create and update (upsert)
const model = await apiClient.post<AdminModel>(
`/admin/api/upstream-providers/${providerId}/models`,
const model = await apiClient.patch<AdminModel>(
`/admin/api/upstream-providers/${providerId}/models/${encodeURIComponent(modelId)}`,
payload
);
return {
+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',
{}
);
}
}