diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 61955035..0f4a4c0b 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -577,7 +577,6 @@ async def _cleanup_enabled_models_once() -> None: 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) @@ -613,7 +612,7 @@ async def _cleanup_enabled_models_once() -> None: def _pricing_matches( - db_pricing: dict, upstream_pricing: dict, tolerance: float = 0.1 + db_pricing: dict, upstream_pricing: dict, tolerance: float = 0.0 ) -> bool: """Check if pricing dictionaries match within tolerance.""" keys_to_compare = [ @@ -626,9 +625,8 @@ def _pricing_matches( ] 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) + 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 diff --git a/routstr/proxy.py b/routstr/proxy.py index 7c63de49..47b98c13 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -3,7 +3,7 @@ from typing import Any from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import Response, StreamingResponse -from sqlmodel import col, select +from sqlmodel import select from .algorithm import create_model_mappings from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key @@ -83,9 +83,7 @@ async def refresh_model_maps() -> None: # Gather database overrides and disabled models async with create_session() as session: - result = await session.exec( - select(ModelRow).where(col(ModelRow.enabled).is_(True)) - ) + result = await session.exec(select(ModelRow).where(ModelRow.enabled)) override_rows = result.all() provider_result = await session.exec(select(UpstreamProviderRow)) @@ -103,7 +101,7 @@ async def refresh_model_maps() -> None: } disabled_result = await session.exec( - select(ModelRow.id).where(col(ModelRow.enabled).is_(False)) + select(ModelRow.id).where(ModelRow.enabled == False) # noqa: E712 ) disabled_model_ids = {row for row in disabled_result.all()} diff --git a/routstr/upstreams/ollama.py b/routstr/upstreams/ollama.py index e999e221..117b95a7 100644 --- a/routstr/upstreams/ollama.py +++ b/routstr/upstreams/ollama.py @@ -121,7 +121,7 @@ class OllamaUpstreamProvider(UpstreamProvider): models_list.append( Model( id=model_name, - name=model_name, + name=model_name.replace(":", " "), created=0, description=description, context_length=context_length, diff --git a/scripts/build-ui.sh b/scripts/build-ui.sh index 557c26db..7a235f01 100755 --- a/scripts/build-ui.sh +++ b/scripts/build-ui.sh @@ -49,6 +49,7 @@ else npm run build fi +rm -rf ../ui_out mkdir -p ../ui_out mv out/* ../ui_out diff --git a/ui/app/layout.tsx b/ui/app/layout.tsx index 1eff9d1c..6c087695 100644 --- a/ui/app/layout.tsx +++ b/ui/app/layout.tsx @@ -34,7 +34,7 @@ export default function RootLayout({ return ( {children} diff --git a/ui/app/model/page.tsx b/ui/app/model/page.tsx index cc13cb7a..edd6b20f 100644 --- a/ui/app/model/page.tsx +++ b/ui/app/model/page.tsx @@ -16,6 +16,7 @@ import { Alert, AlertDescription } from '@/components/ui/alert'; import { Badge } from '@/components/ui/badge'; import { useMemo, useState } from 'react'; import type { Model } from '@/lib/api/schemas/models'; +import { groupAndSortModelsByProvider } from '@/lib/utils/modelSort'; export default function ModelsPage() { const [filteredModels, setFilteredModels] = useState([]); @@ -34,15 +35,7 @@ export default function ModelsPage() { const groupedModels = useMemo(() => { if (!models) return {}; - - return models.reduce>((acc, model) => { - const provider = model.provider; - if (!acc[provider]) { - acc[provider] = []; - } - acc[provider].push(model); - return acc; - }, {}); + return groupAndSortModelsByProvider(models); }, [models]); const groupDataMap = useMemo(() => { @@ -52,7 +45,9 @@ export default function ModelsPage() { const providerInfo = useMemo(() => { return Object.entries(groupedModels).map(([provider, providerModels]) => { const groupData = groupDataMap.get(provider); - const activeModels = providerModels.filter((m) => !m.soft_deleted).length; + const activeModels = providerModels.filter( + (m) => m.isEnabled && !m.soft_deleted + ).length; const totalModels = providerModels.length; return { @@ -156,7 +151,10 @@ export default function ModelsPage() { models={models} onFilteredModelsChange={setFilteredModels} /> - + @@ -183,7 +181,7 @@ export default function ModelsPage() { (m) => m.soft_deleted ).length }{' '} - soft deleted + disabled )} {groupData?.group_url && ( @@ -199,6 +197,7 @@ export default function ModelsPage() { filterProvider={provider} groupData={groupData} showProviderActions={true} + showDeleteAllButton={false} /> diff --git a/ui/app/providers/page.tsx b/ui/app/providers/page.tsx index 49411bb0..53ca72c4 100644 --- a/ui/app/providers/page.tsx +++ b/ui/app/providers/page.tsx @@ -277,7 +277,11 @@ export default function ProvidersPage() { } placeholder='https://api.example.com/v1' disabled={hasFixedBaseUrl(formData.provider_type)} - className={hasFixedBaseUrl(formData.provider_type) ? 'cursor-not-allowed opacity-60' : ''} + className={ + hasFixedBaseUrl(formData.provider_type) + ? 'cursor-not-allowed opacity-60' + : '' + } />
@@ -453,7 +457,8 @@ export default function ProvidersPage() {
{providerModels.db_models.length === 0 ? (
- No models configured. Add custom models to use this provider. + No models configured. Add custom models to + use this provider.
) : (
@@ -495,7 +500,10 @@ export default function ProvidersPage() {
) : ( // Has provided models - show tabs - + Provided Models - Provided + + Provided + {providerModels.db_models.length > 0 && (
- Custom models override or extend the provider's catalog. + Custom models override or extend the + provider's catalog.
)} {providerModels.db_models.length === 0 ? ( @@ -543,39 +554,42 @@ export default function ProvidersPage() {
) : (
- {providerModels.db_models.map((model) => ( -
-
-
- - {model.id} - - - {model.enabled - ? 'Enabled' - : 'Disabled'} - + {providerModels.db_models.map( + (model) => ( +
+
+
+ + {model.id} + + + {model.enabled + ? 'Enabled' + : 'Disabled'} + +
+
+ {model.description || + model.name} +
-
- {model.description || model.name} +
+ {model.context_length?.toLocaleString()}{' '} + tokens
-
- {model.context_length?.toLocaleString()}{' '} - tokens -
-
- ))} + ) + )}
)} @@ -583,9 +597,11 @@ export default function ProvidersPage() { value='provided' className='mt-4 space-y-2' > - {providerModels.remote_models.length > 0 && ( + {providerModels.remote_models.length > + 0 && (
- Models automatically discovered from the provider's catalog. + Models automatically discovered from the + provider's catalog.
)}
@@ -670,7 +686,11 @@ export default function ProvidersPage() { } placeholder='https://api.example.com/v1' disabled={hasFixedBaseUrl(formData.provider_type)} - className={hasFixedBaseUrl(formData.provider_type) ? 'cursor-not-allowed opacity-60' : ''} + className={ + hasFixedBaseUrl(formData.provider_type) + ? 'cursor-not-allowed opacity-60' + : '' + } />
diff --git a/ui/components/ModelSelector.tsx b/ui/components/ModelSelector.tsx index 5b03e350..914cd850 100644 --- a/ui/components/ModelSelector.tsx +++ b/ui/components/ModelSelector.tsx @@ -59,18 +59,24 @@ import { } from 'lucide-react'; import { toast } from 'sonner'; import { cn } from '@/lib/utils'; +import { + sortModels, + groupAndSortModelsByProvider, +} from '@/lib/utils/modelSort'; interface ModelSelectorProps { filterProvider?: string; groupData?: ModelGroup; showProviderActions?: boolean; filteredModels?: Model[]; + showDeleteAllButton?: boolean; } export function ModelSelector({ filterProvider, groupData, filteredModels: propFilteredModels, + showDeleteAllButton = false, }: ModelSelectorProps) { const [selectedModelId, setSelectedModelId] = useState(''); const [, setHoveredModelId] = useState(null); @@ -387,14 +393,10 @@ export function ModelSelector({ providerIdNum, model.id ); - await AdminService.updateProviderModel( - providerIdNum, - model.full_name, - { - ...existingModel, - enabled: true, - } - ); + await AdminService.updateProviderModel(providerIdNum, model.id, { + ...existingModel, + enabled: true, + }); totalEnabled++; } catch (error) { console.error(`Failed to enable model ${model.full_name}:`, error); @@ -426,20 +428,13 @@ export function ModelSelector({ : providerFilteredModels; if (filterProvider) { - // If filtering by provider, return single group - return { [filterProvider]: modelsToGroup }; + const sortedModels = sortModels(modelsToGroup); + return { [filterProvider]: sortedModels }; } if (!modelsToGroup) return {}; - return modelsToGroup.reduce>((acc, model) => { - const provider = model.provider; - if (!acc[provider]) { - acc[provider] = []; - } - acc[provider].push(model); - return acc; - }, {}); + return groupAndSortModelsByProvider(modelsToGroup); }, [ providerFilteredModels, filteredModels, @@ -828,13 +823,15 @@ export function ModelSelector({ Deselect All - + {showDeleteAllButton && ( + + )} {/* Model Management Actions @@ -1181,7 +1178,7 @@ export function ModelSelector({ {model.soft_deleted && ( - Deleted + Disabled )}
diff --git a/ui/components/detailed-wallet-balance.tsx b/ui/components/detailed-wallet-balance.tsx index e5dd69b4..9979b0f4 100644 --- a/ui/components/detailed-wallet-balance.tsx +++ b/ui/components/detailed-wallet-balance.tsx @@ -203,7 +203,8 @@ export function DetailedWalletBalance({
0 && + !detail.error && + ownerMsat > 0 && 'font-semibold text-green-600' )} > @@ -250,7 +251,8 @@ export function DetailedWalletBalance({
0 && + !detail.error && + ownerMsat > 0 && 'font-semibold text-green-600' )} > @@ -283,4 +285,4 @@ export function DetailedWalletBalance({ /> ); -} \ No newline at end of file +} diff --git a/ui/components/temporary-balances.tsx b/ui/components/temporary-balances.tsx index 656f3aff..c33d0d43 100644 --- a/ui/components/temporary-balances.tsx +++ b/ui/components/temporary-balances.tsx @@ -11,10 +11,7 @@ import { DollarSign, Activity, } from 'lucide-react'; -import { - AdminService, - TemporaryBalance, -} from '@/lib/api/services/admin'; +import { AdminService, TemporaryBalance } from '@/lib/api/services/admin'; import { Card, CardContent, @@ -243,16 +240,16 @@ export function TemporaryBalances({
Balance
-
- {formatBalance(balance.balance)} +
+ {formatBalance(balance.balance)}
Spent
-
- {formatBalance(balance.total_spent)} +
+ {formatBalance(balance.total_spent)}
diff --git a/ui/lib/currency.ts b/ui/lib/currency.ts index bd5fb3d6..9f9ab595 100644 --- a/ui/lib/currency.ts +++ b/ui/lib/currency.ts @@ -43,4 +43,3 @@ export function formatFromMsat( }); return formatter.format(usd); } - diff --git a/ui/lib/exchange-rate.ts b/ui/lib/exchange-rate.ts index a4131611..7605bb7a 100644 --- a/ui/lib/exchange-rate.ts +++ b/ui/lib/exchange-rate.ts @@ -51,4 +51,3 @@ export async function fetchBtcUsdPrice(): Promise { export function btcToSatsRate(btcUsdPrice: number): number { return btcUsdPrice / 100_000_000; } - diff --git a/ui/lib/types/units.ts b/ui/lib/types/units.ts index 162e8613..0b7f776d 100644 --- a/ui/lib/types/units.ts +++ b/ui/lib/types/units.ts @@ -14,4 +14,3 @@ export function getDisplayUnitLabel(unit: DisplayUnit): string { return unit; } } - diff --git a/ui/lib/utils/modelSort.ts b/ui/lib/utils/modelSort.ts new file mode 100644 index 00000000..932cccc1 --- /dev/null +++ b/ui/lib/utils/modelSort.ts @@ -0,0 +1,36 @@ +import { type Model } from '@/lib/api/schemas/models'; + +export function sortModelsByStatus(a: Model, b: Model): number { + if (a.isEnabled && !b.isEnabled) return -1; + if (!a.isEnabled && b.isEnabled) return 1; + + if (a.isEnabled === b.isEnabled) { + if (!a.soft_deleted && b.soft_deleted) return -1; + if (a.soft_deleted && !b.soft_deleted) return 1; + } + + return 0; +} + +export function sortModels(models: Model[]): Model[] { + return [...models].sort(sortModelsByStatus); +} + +export function groupAndSortModelsByProvider( + models: Model[] +): Record { + const grouped = models.reduce>((acc, model) => { + const provider = model.provider; + if (!acc[provider]) { + acc[provider] = []; + } + acc[provider].push(model); + return acc; + }, {}); + + Object.keys(grouped).forEach((provider) => { + grouped[provider].sort(sortModelsByStatus); + }); + + return grouped; +}