unsloth/studio/backend/models/providers.py
Roland Tannous 4656e63374 studio: add external provider support for chat inference
Adds the ability to connect to OpenAI, Mistral, Google, Cohere, Together,
Fireworks, and Perplexity from the Studio chat interface.

- Provider configs stored in SQLite (no API keys persisted)
- RSA-2048 key pair generated at startup for client-side key encryption
- httpx proxy client streams SSE responses in OpenAI-compatible format
- New /api/providers routes: registry, CRUD, test, models
- /v1/chat/completions routes to external provider when provider fields present
- Integration test suite covering CRUD, connection, model listing, and inference
- Frontend spec doc with full API contract
2026-03-29 18:12:45 +00:00

106 lines
4.6 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""
Pydantic schemas for the external LLM providers API.
"""
from typing import Optional
from pydantic import BaseModel, Field
# ── Registry (static provider info) ───────────────────────────────
class ProviderRegistryEntry(BaseModel):
"""A supported provider type with its default configuration."""
provider_type: str = Field(..., description="Provider identifier (e.g. 'openai', 'mistral')")
display_name: str = Field(..., description="Human-readable provider name")
base_url: str = Field(..., description="Default API base URL")
default_models: list[str] = Field(
default_factory=list, description="Well-known model IDs for this provider"
)
supports_streaming: bool = Field(True, description="Whether this provider supports SSE streaming")
supports_vision: bool = Field(False, description="Whether this provider supports vision/image input")
supports_tool_calling: bool = Field(False, description="Whether this provider supports tool/function calling")
# ── Provider config CRUD ──────────────────────────────────────────
class ProviderCreate(BaseModel):
"""Request to create a saved provider configuration."""
provider_type: str = Field(..., description="Provider type from the registry")
display_name: str = Field(..., description="User-chosen label (e.g. 'My OpenAI Key')")
base_url: Optional[str] = Field(
None,
description="Custom base URL (overrides registry default). Omit to use the default.",
)
class ProviderUpdate(BaseModel):
"""Request to update a saved provider configuration."""
display_name: Optional[str] = Field(None, description="New display name")
base_url: Optional[str] = Field(None, description="New base URL")
is_enabled: Optional[bool] = Field(None, description="Enable or disable this provider")
class ProviderResponse(BaseModel):
"""A saved provider configuration (returned by list/get endpoints)."""
id: str = Field(..., description="Unique provider config ID")
provider_type: str = Field(..., description="Provider type (e.g. 'openai')")
display_name: str = Field(..., description="User-chosen label")
base_url: str = Field(..., description="API base URL")
is_enabled: bool = Field(True, description="Whether this provider is enabled")
created_at: str = Field(..., description="ISO 8601 creation timestamp")
updated_at: str = Field(..., description="ISO 8601 last-update timestamp")
# ── Model listing ─────────────────────────────────────────────────
class ProviderModelInfo(BaseModel):
"""A model available from an external provider."""
id: str = Field(..., description="Model ID as expected by the provider API")
display_name: str = Field("", description="Human-readable model name")
context_length: Optional[int] = Field(None, description="Maximum context length in tokens")
owned_by: Optional[str] = Field(None, description="Model owner/organization")
class ProviderModelsRequest(BaseModel):
"""Request to list models from an external provider."""
provider_type: str = Field(..., description="Provider type from the registry")
encrypted_api_key: str = Field(..., description="RSA-encrypted, base64-encoded API key")
base_url: Optional[str] = Field(
None, description="Custom base URL (overrides registry default)"
)
# ── Connection testing ────────────────────────────────────────────
class ProviderTestRequest(BaseModel):
"""Request to test connectivity to an external provider."""
provider_type: str = Field(..., description="Provider type from the registry")
encrypted_api_key: str = Field(..., description="RSA-encrypted, base64-encoded API key")
base_url: Optional[str] = Field(
None, description="Custom base URL (overrides registry default)"
)
class ProviderTestResult(BaseModel):
"""Result of a provider connectivity test."""
success: bool = Field(..., description="Whether the test succeeded")
message: str = Field(..., description="Human-readable result message")
models_count: Optional[int] = Field(
None, description="Number of models found (if test succeeded)"
)