mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
better model update
This commit is contained in:
@@ -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:
|
||||
@@ -87,6 +89,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
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())
|
||||
nip91_task = asyncio.create_task(announce_provider())
|
||||
@@ -115,6 +118,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()
|
||||
|
||||
@@ -132,6 +137,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)
|
||||
|
||||
|
||||
@@ -526,6 +526,116 @@ 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
|
||||
print(db_model.id)
|
||||
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.1
|
||||
) -> 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 = float(db_pricing.get(key, 0.0)) * 1000000
|
||||
upstream_val = float(upstream_pricing.get(key, 0.0)) * 1000000
|
||||
print(db_val - upstream_val)
|
||||
|
||||
if abs(db_val - upstream_val) > tolerance:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
|
||||
async def refresh_models_periodically() -> None:
|
||||
"""Background task: periodically fetch OpenRouter models and insert new ones.
|
||||
|
||||
|
||||
@@ -121,7 +121,7 @@ export function EditModelForm({
|
||||
|
||||
const adminModel = await AdminService.getProviderModel(
|
||||
providerId,
|
||||
model.full_name
|
||||
model.id
|
||||
);
|
||||
|
||||
setAdminModelData(adminModel as AdminModelData);
|
||||
|
||||
@@ -139,7 +139,7 @@ export function ModelSelector({
|
||||
if (providerId === 'unknown') {
|
||||
continue;
|
||||
}
|
||||
const modelFullNames = providerModels.map((m) => m.full_name);
|
||||
const modelFullNames = providerModels.map((m) => m.id);
|
||||
const result = await AdminService.deleteModels(
|
||||
modelFullNames,
|
||||
providerId
|
||||
@@ -191,22 +191,18 @@ export function ModelSelector({
|
||||
try {
|
||||
const existingModel = await AdminService.getProviderModel(
|
||||
providerIdNum,
|
||||
model.full_name
|
||||
);
|
||||
await AdminService.updateProviderModel(
|
||||
providerIdNum,
|
||||
model.full_name,
|
||||
{
|
||||
...existingModel,
|
||||
enabled: false,
|
||||
}
|
||||
model.id
|
||||
);
|
||||
await AdminService.updateProviderModel(providerIdNum, model.id, {
|
||||
...existingModel,
|
||||
enabled: false,
|
||||
});
|
||||
totalDisabled++;
|
||||
} catch (fetchError: unknown) {
|
||||
const error = fetchError as { message?: string; status?: number };
|
||||
if (error.message?.includes('404') || error.status === 404) {
|
||||
const newOverride = {
|
||||
id: model.full_name,
|
||||
id: model.id,
|
||||
name: model.name,
|
||||
description: model.description || '',
|
||||
created: Math.floor(Date.now() / 1000),
|
||||
@@ -340,10 +336,10 @@ export function ModelSelector({
|
||||
const providerId = parseInt(model.provider_id);
|
||||
const existingModel = await AdminService.getProviderModel(
|
||||
providerId,
|
||||
model.full_name
|
||||
model.id
|
||||
);
|
||||
|
||||
await AdminService.updateProviderModel(providerId, model.full_name, {
|
||||
await AdminService.updateProviderModel(providerId, model.id, {
|
||||
...existingModel,
|
||||
enabled: true,
|
||||
});
|
||||
@@ -389,7 +385,7 @@ export function ModelSelector({
|
||||
try {
|
||||
const existingModel = await AdminService.getProviderModel(
|
||||
providerIdNum,
|
||||
model.full_name
|
||||
model.id
|
||||
);
|
||||
await AdminService.updateProviderModel(
|
||||
providerIdNum,
|
||||
@@ -747,7 +743,7 @@ export function ModelSelector({
|
||||
try {
|
||||
const existingModel = await AdminService.getProviderModel(
|
||||
providerId,
|
||||
model.full_name
|
||||
model.id
|
||||
);
|
||||
|
||||
await AdminService.updateProviderModel(providerId, model.full_name, {
|
||||
@@ -759,7 +755,7 @@ export function ModelSelector({
|
||||
const error = fetchError as { message?: string; status?: number };
|
||||
if (error.message?.includes('404') || error.status === 404) {
|
||||
const newOverride = {
|
||||
id: model.full_name,
|
||||
id: model.id,
|
||||
name: model.name,
|
||||
description: model.description || '',
|
||||
created: Math.floor(Date.now() / 1000),
|
||||
|
||||
@@ -497,7 +497,7 @@ export class AdminService {
|
||||
internal_reasoning: 0,
|
||||
};
|
||||
|
||||
const modelId = (data.id as string) || (data.full_name as string);
|
||||
const modelId = data.id as string;
|
||||
|
||||
const payload = {
|
||||
model_id: modelId,
|
||||
|
||||
Reference in New Issue
Block a user