unsloth/studio/backend/core/inference/key_exchange.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

71 lines
2.2 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
"""
RSA key pair for encrypting API keys in transit.
The frontend encrypts API keys with the server's public key before
including them in requests. The backend decrypts with its private key
before forwarding to external providers.
The key pair is generated at server startup and lives only in memory —
it is regenerated on each restart. The frontend fetches the public key
via GET /api/providers/public-key on load.
"""
import base64
import logging
from cryptography.hazmat.primitives.asymmetric import rsa, padding
from cryptography.hazmat.primitives import serialization, hashes
logger = logging.getLogger(__name__)
_private_key: rsa.RSAPrivateKey | None = None
_public_key_pem: str | None = None
def init_key_pair() -> None:
"""Generate an RSA-2048 key pair. Called once at server startup."""
global _private_key, _public_key_pem
_private_key = rsa.generate_private_key(
public_exponent=65537,
key_size=2048,
)
_public_key_pem = _private_key.public_key().public_bytes(
serialization.Encoding.PEM,
serialization.PublicFormat.SubjectPublicKeyInfo,
).decode("utf-8")
logger.info("RSA key pair generated for API key encryption")
def get_public_key_pem() -> str:
"""Return the PEM-encoded public key for the frontend."""
if _public_key_pem is None:
raise RuntimeError("Key pair not initialized. Call init_key_pair() first.")
return _public_key_pem
def decrypt_api_key(encrypted_b64: str) -> str:
"""
Decrypt an API key that was encrypted with the public key.
Args:
encrypted_b64: Base64-encoded RSA-OAEP ciphertext.
Returns:
The plaintext API key string.
"""
if _private_key is None:
raise RuntimeError("Key pair not initialized. Call init_key_pair() first.")
ciphertext = base64.b64decode(encrypted_b64)
plaintext = _private_key.decrypt(
ciphertext,
padding.OAEP(
mgf=padding.MGF1(algorithm=hashes.SHA256()),
algorithm=hashes.SHA256(),
label=None,
),
)
return plaintext.decode("utf-8")