From b94b175c81473e40838907e4c39b4086760864ec Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 3 Oct 2026 01:24:42 +0200 Subject: [PATCH] revert: keep /api/models/test and Basic Testing until #760 lands --- routstr/payment/models.py | 112 ++++- .../test_model_test_endpoint_security.py | 224 +++++++++ ui/components/model-tester.tsx | 450 ++++++++++++++++++ ui/components/models-page.tsx | 167 ++++--- ui/lib/api/schemas/models.ts | 26 + ui/lib/api/services/models.ts | 25 + 6 files changed, 949 insertions(+), 55 deletions(-) create mode 100644 tests/integration/test_model_test_endpoint_security.py create mode 100644 ui/components/model-tester.tsx diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 916dfe87..8d7788a1 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -3,11 +3,12 @@ import json import random import httpx -from fastapi import APIRouter, Depends, HTTPException +from fastapi import APIRouter, Depends, HTTPException, Request +from pydantic import BaseModel as V2BaseModel from pydantic.v1 import BaseModel, validator from sqlmodel.ext.asyncio.session import AsyncSession -from ..core.db import ModelRow, get_session +from ..core.db import ModelRow, UpstreamProviderRow, get_session from ..core.logging import get_logger from ..core.settings import settings from .price import sats_usd_price @@ -17,6 +18,24 @@ logger = get_logger(__name__) models_router = APIRouter() +_MODEL_TEST_ENDPOINT_PATHS = { + "chat-completions": "chat/completions", + "completions": "completions", + "embeddings": "embeddings", + "responses": "responses", +} + +# Cap the caller-supplied test payload to avoid forwarding oversized bodies +# upstream on the operator's credentials. +_MODEL_TEST_MAX_REQUEST_BYTES = 64 * 1024 + + +async def _require_admin_api(request: Request) -> None: + """Require admin auth without creating an import-time cycle with core.admin.""" + from ..core.admin import require_admin_api + + await require_admin_api(request) + class Architecture(BaseModel): modality: str @@ -630,6 +649,95 @@ async def update_sats_pricing() -> None: logger.error(f"Error updating sats pricing: {e}") +class ModelTestRequest(V2BaseModel): + model_id: str + endpoint_type: str + request_data: dict + + +@models_router.post("/api/models/test", dependencies=[Depends(_require_admin_api)]) +async def test_model( + payload: ModelTestRequest, + session: AsyncSession = Depends(get_session), +) -> dict: + """Test a model by sending a request through its configured upstream provider.""" + from sqlmodel import select + + result = await session.execute( + select(ModelRow).where(ModelRow.id == payload.model_id) + ) + model_row = result.scalars().first() + + if not model_row: + return { + "success": False, + "error": f"Model '{payload.model_id}' not found in database", + "status_code": 404, + } + + provider = await session.get(UpstreamProviderRow, model_row.upstream_provider_id) + if not provider: + return { + "success": False, + "error": "Upstream provider not found", + "status_code": 404, + } + + endpoint_path = _MODEL_TEST_ENDPOINT_PATHS.get(payload.endpoint_type) + if endpoint_path is None: + raise HTTPException(status_code=400, detail="Unsupported endpoint_type") + + actual_model_id = model_row.forwarded_model_id or model_row.id + request_data = dict(payload.request_data) + request_data["model"] = actual_model_id + + try: + request_size = len(json.dumps(request_data).encode("utf-8")) + except (TypeError, ValueError): + raise HTTPException(status_code=400, detail="Invalid request_data") + if request_size > _MODEL_TEST_MAX_REQUEST_BYTES: + raise HTTPException(status_code=413, detail="request_data too large") + + base_url = provider.base_url.rstrip("/") + url = f"{base_url}/{endpoint_path}" + + logger.info( + "admin model test", + extra={ + "model_id": payload.model_id, + "forwarded_model_id": actual_model_id, + "endpoint_type": payload.endpoint_type, + "upstream_provider_id": model_row.upstream_provider_id, + "request_bytes": request_size, + }, + ) + + headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {provider.api_key}", + } + + try: + async with httpx.AsyncClient(timeout=30.0) as client: + response = await client.post(url, json=request_data, headers=headers) + try: + response_data = response.json() + except Exception: + response_data = {"raw": response.text} + + return { + "success": response.status_code < 400, + "data": response_data, + "status_code": response.status_code, + } + except Exception as e: + return { + "success": False, + "error": str(e), + "status_code": 500, + } + + @models_router.get("/v1/models/paths") @models_router.get("/v1/models/paths/", include_in_schema=False) async def model_paths() -> dict: diff --git a/tests/integration/test_model_test_endpoint_security.py b/tests/integration/test_model_test_endpoint_security.py new file mode 100644 index 00000000..dca7e9fa --- /dev/null +++ b/tests/integration/test_model_test_endpoint_security.py @@ -0,0 +1,224 @@ +import time +from types import TracebackType +from typing import Any +from unittest.mock import patch + +import pytest +from httpx import AsyncClient + +from routstr.core.admin import admin_sessions +from routstr.core.db import ModelRow, UpstreamProviderRow + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_model_test_endpoint_requires_admin_auth( + integration_client: AsyncClient, +) -> None: + with patch("httpx.AsyncClient") as mock_async_client: + response = await integration_client.post( + "/api/models/test", + json={ + "model_id": "model-a", + "endpoint_type": "chat-completions", + "request_data": {"messages": []}, + }, + ) + + assert response.status_code == 403 + mock_async_client.assert_not_called() + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_model_test_endpoint_rejects_unsupported_endpoint_type( + integration_client: AsyncClient, + integration_session: Any, +) -> None: + admin_token = "test-admin-token-model-test" + admin_sessions[admin_token] = int(time.time()) + 3600 + integration_client.headers["Authorization"] = f"Bearer {admin_token}" + + provider = UpstreamProviderRow( + provider_type="custom", + base_url="https://api.example.com/v1", + api_key="sk-upstream-test", + enabled=True, + provider_fee=1.01, + ) + integration_session.add(provider) + await integration_session.commit() + await integration_session.refresh(provider) + assert provider.id is not None + + model = ModelRow( + id="model-a", + name="Model A", + created=1, + description="desc", + context_length=100, + architecture='{"modality": "text", "input_modalities": ["text"], "output_modalities": ["text"], "tokenizer": "tiktoken", "instruct_type": "chat"}', + pricing='{"prompt": 1.0, "completion": 1.0}', + upstream_provider_id=provider.id, + enabled=True, + ) + integration_session.add(model) + await integration_session.commit() + + try: + with patch("httpx.AsyncClient") as mock_async_client: + response = await integration_client.post( + "/api/models/test", + json={ + "model_id": "model-a", + "endpoint_type": "../../abuse", + "request_data": {"messages": []}, + }, + ) + + assert response.status_code == 400 + assert response.json()["detail"] == "Unsupported endpoint_type" + mock_async_client.assert_not_called() + finally: + admin_sessions.pop(admin_token, None) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_model_test_endpoint_rejects_oversized_request_data( + integration_client: AsyncClient, + integration_session: Any, +) -> None: + admin_token = "test-admin-token-model-test-oversized" + admin_sessions[admin_token] = int(time.time()) + 3600 + integration_client.headers["Authorization"] = f"Bearer {admin_token}" + + provider = UpstreamProviderRow( + provider_type="custom", + base_url="https://api.example.com/v1", + api_key="sk-upstream-test", + enabled=True, + provider_fee=1.01, + ) + integration_session.add(provider) + await integration_session.commit() + await integration_session.refresh(provider) + assert provider.id is not None + + model = ModelRow( + id="model-a", + name="Model A", + created=1, + description="desc", + context_length=100, + architecture='{"modality": "text", "input_modalities": ["text"], "output_modalities": ["text"], "tokenizer": "tiktoken", "instruct_type": "chat"}', + pricing='{"prompt": 1.0, "completion": 1.0}', + upstream_provider_id=provider.id, + enabled=True, + ) + integration_session.add(model) + await integration_session.commit() + + oversized = "x" * (64 * 1024 + 1) + + try: + with patch("httpx.AsyncClient") as mock_async_client: + response = await integration_client.post( + "/api/models/test", + json={ + "model_id": "model-a", + "endpoint_type": "chat-completions", + "request_data": {"blob": oversized}, + }, + ) + + assert response.status_code == 413 + assert response.json()["detail"] == "request_data too large" + mock_async_client.assert_not_called() + finally: + admin_sessions.pop(admin_token, None) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_model_test_endpoint_admin_uses_allowed_upstream_path( + integration_client: AsyncClient, + integration_session: Any, +) -> None: + admin_token = "test-admin-token-model-test-success" + admin_sessions[admin_token] = int(time.time()) + 3600 + integration_client.headers["Authorization"] = f"Bearer {admin_token}" + + provider = UpstreamProviderRow( + provider_type="custom", + base_url="https://api.example.com/v1", + api_key="sk-upstream-test", + enabled=True, + provider_fee=1.01, + ) + integration_session.add(provider) + await integration_session.commit() + await integration_session.refresh(provider) + assert provider.id is not None + + model = ModelRow( + id="model-a", + name="Model A", + created=1, + description="desc", + context_length=100, + architecture='{"modality": "text", "input_modalities": ["text"], "output_modalities": ["text"], "tokenizer": "tiktoken", "instruct_type": "chat"}', + pricing='{"prompt": 1.0, "completion": 1.0}', + upstream_provider_id=provider.id, + enabled=True, + forwarded_model_id="upstream-model-a", + ) + integration_session.add(model) + await integration_session.commit() + + class MockResponse: + status_code = 200 + text = '{"ok": true}' + + def json(self) -> dict[str, bool]: + return {"ok": True} + + class MockAsyncClient: + async def __aenter__(self) -> "MockAsyncClient": + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + tb: TracebackType | None, + ) -> None: + return None + + async def post( + self, url: str, json: dict[str, Any], headers: dict[str, str] + ) -> MockResponse: + assert url == "https://api.example.com/v1/chat/completions" + assert json["model"] == "upstream-model-a" + assert headers["Authorization"] == "Bearer sk-upstream-test" + return MockResponse() + + try: + with patch("httpx.AsyncClient", return_value=MockAsyncClient()): + response = await integration_client.post( + "/api/models/test", + json={ + "model_id": "model-a", + "endpoint_type": "chat-completions", + "request_data": {"messages": []}, + }, + ) + + assert response.status_code == 200 + assert response.json() == { + "success": True, + "data": {"ok": True}, + "status_code": 200, + } + finally: + admin_sessions.pop(admin_token, None) diff --git a/ui/components/model-tester.tsx b/ui/components/model-tester.tsx new file mode 100644 index 00000000..3ecd43b3 --- /dev/null +++ b/ui/components/model-tester.tsx @@ -0,0 +1,450 @@ +'use client'; + +import React, { useState } from 'react'; +import { useMutation, useQuery } from '@tanstack/react-query'; +import { type Model } from '@/lib/api/schemas/models'; +import { ModelService } from '@/lib/api/services/models'; +import { Button } from '@/components/ui/button'; +import { Textarea } from '@/components/ui/textarea'; +import { Input } from '@/components/ui/input'; +import { Label } from '@/components/ui/label'; +import { + Card, + CardContent, + CardDescription, + CardHeader, + CardTitle, +} from '@/components/ui/card'; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from '@/components/ui/select'; +import { Badge } from '@/components/ui/badge'; +import { Alert, AlertDescription } from '@/components/ui/alert'; +import { + Loader2, + Send, + CheckCircle, + XCircle, + Info, + Key, + Globe, +} from 'lucide-react'; +import { toast } from 'sonner'; + +interface ModelTesterProps { + models: Model[]; +} + +interface ChatCompletionRequest { + model: string; + messages: { + role: 'system' | 'user' | 'assistant'; + content: string; + }[]; + max_tokens?: number; + temperature?: number; +} + +interface ChatCompletionResponse { + id: string; + object: string; + created: number; + model: string; + choices: { + index: number; + message: { + role: string; + content: string; + }; + finish_reason: string; + }[]; + usage?: { + prompt_tokens: number; + completion_tokens: number; + total_tokens: number; + }; +} + +const DEFAULT_SYSTEM_MESSAGE = + 'You are a helpful assistant. Please respond concisely.'; +const DEFAULT_USER_MESSAGE = + 'Hello! Can you tell me what model you are and confirm that you are working correctly?'; + +export function ModelTester({ models }: ModelTesterProps) { + const [selectedModelId, setSelectedModelId] = useState(''); + const [systemMessage, setSystemMessage] = useState(DEFAULT_SYSTEM_MESSAGE); + const [userMessage, setUserMessage] = useState(DEFAULT_USER_MESSAGE); + const [maxTokens, setMaxTokens] = useState(150); + const [temperature, setTemperature] = useState(0.7); + const [response, setResponse] = useState(null); + const [error, setError] = useState(null); + + // Fetch model groups for API key resolution + const { data: groups = [] } = useQuery({ + queryKey: ['model-groups'], + queryFn: () => ModelService.getModelGroups(), + refetchOnWindowFocus: false, + }); + + const selectedModel = models.find((model) => model.id === selectedModelId); + + // Get effective API key and endpoint URL for the selected model + const getModelCredentials = (model: Model) => { + const group = groups.find((g) => g.provider === model.provider); + + // Determine API key (individual takes precedence over group) + const apiKey = model.api_key || group?.group_api_key; + + // Determine endpoint URL + let endpointUrl = model.url; + + // If model URL is relative and group has a base URL, combine them + if (model.url.startsWith('/') && group?.group_url) { + endpointUrl = `${group.group_url.replace(/\/$/, '')}${model.url}`; + } + + // Ensure the URL ends with /chat/completions for chat models + if ( + model.modelType === 'text' && + !endpointUrl.includes('/chat/completions') + ) { + endpointUrl = endpointUrl.replace(/\/$/, '') + '/chat/completions'; + } + + return { + apiKey, + endpointUrl, + group, + }; + }; + + const testModelMutation = useMutation({ + mutationFn: async (request: ChatCompletionRequest) => { + if (!selectedModel) { + throw new Error('No model selected'); + } + + setError(null); + setResponse(null); + + try { + const response = await ModelService.testModel( + selectedModel.id, + 'chat-completions', + request + ); + + if (!response.success) { + throw new Error(response.error || 'Test failed'); + } + + return response.data as ChatCompletionResponse; + } catch (err: unknown) { + console.error('Model test error via proxy:', err); + + const errorMessage = + err instanceof Error ? err.message : 'Failed to test model via proxy'; + throw new Error(errorMessage); + } + }, + onSuccess: (data) => { + setResponse(data); + toast.success('Model test completed successfully!'); + }, + onError: (err: Error) => { + const errorMessage = err?.message || 'Unknown error occurred'; + setError(errorMessage); + toast.error(`Model test failed: ${errorMessage}`); + }, + }); + + const handleTest = async () => { + if (!selectedModel) { + toast.error('Please select a model to test'); + return; + } + + if (!userMessage.trim()) { + toast.error('Please enter a test message'); + return; + } + + const messages = []; + + if (systemMessage.trim()) { + messages.push({ + role: 'system' as const, + content: systemMessage.trim(), + }); + } + + messages.push({ + role: 'user' as const, + content: userMessage.trim(), + }); + + const request: ChatCompletionRequest = { + model: selectedModel.name, + messages, + max_tokens: maxTokens, + temperature: temperature, + }; + + testModelMutation.mutate(request); + }; + + const enabledModels = Array.from( + new Map( + models.filter((model) => model.isEnabled).map((m) => [m.id, m]) + ).values() + ); + const credentials = selectedModel ? getModelCredentials(selectedModel) : null; + + return ( + + + Model Credential Tester + + Test model functionality by sending chat completion requests through + the secure proxy (resolves CORS and network issues) + + + + {/* Model Selection */} +
+ + + + {selectedModel && credentials && ( +
+
+ + + Endpoint: {credentials.endpointUrl} + +
+
+ + + API Key:{' '} + {credentials.apiKey + ? `${credentials.apiKey.substring(0, 8)}...` + : 'Not configured'} + + + {selectedModel.api_key_type || 'Unknown'} + +
+
+ + Provider: {selectedModel.provider} + +
+
+ + Type: {selectedModel.modelType} + +
+ {selectedModel.contextLength && ( +
+ + Context Length:{' '} + {selectedModel.contextLength.toLocaleString()} + +
+ )} + {!credentials.apiKey && ( + + + + No API key configured for this model. Testing may still work + if the model is free or if authentication is handled + elsewhere. For models requiring authentication, please add + an API key to the model or its provider group. + + + )} +
+ )} +
+ + {/* Test Parameters */} +
+
+ + setMaxTokens(parseInt(e.target.value) || 150)} + name='max_tokens' + /> +
+
+ + + setTemperature(parseFloat(e.target.value) || 0.7) + } + name='temperature' + /> +
+
+ + {/* System Message */} +
+ +