9b8483e797
This commit introduces: - AES-256-GCM encryption for LLM provider API keys in the database. - HMAC-SHA256 signed session tokens with activity-based refresh logic. - Standardized frontend XSS protection using a global escapeHtml utility. - Hardened security headers and request body size limits. - Improved database integrity with foreign key enforcement and atomic transactions. - Integration tests for the full encrypted key storage and proxy usage lifecycle.
329 lines
12 KiB
Rust
329 lines
12 KiB
Rust
//! LLM Proxy Library
|
|
//!
|
|
//! This library provides the core functionality for the LLM proxy gateway,
|
|
//! including provider integration, token tracking, and API endpoints.
|
|
|
|
pub mod auth;
|
|
pub mod client;
|
|
pub mod config;
|
|
pub mod dashboard;
|
|
pub mod database;
|
|
pub mod errors;
|
|
pub mod logging;
|
|
pub mod models;
|
|
pub mod multimodal;
|
|
pub mod providers;
|
|
pub mod rate_limiting;
|
|
pub mod server;
|
|
pub mod state;
|
|
pub mod utils;
|
|
|
|
// Re-exports for convenience
|
|
pub use auth::{AuthenticatedClient, validate_token};
|
|
pub use config::{
|
|
AppConfig, DatabaseConfig, DeepSeekConfig, GeminiConfig, GrokConfig, ModelMappingConfig, ModelPricing,
|
|
OllamaConfig, OpenAIConfig, PricingConfig, ProviderConfig, ServerConfig,
|
|
};
|
|
pub use database::{DbPool, init as init_db, test_connection};
|
|
pub use errors::AppError;
|
|
pub use logging::{LoggingContext, RequestLog, RequestLogger};
|
|
pub use models::{
|
|
ChatChoice, ChatCompletionRequest, ChatCompletionResponse, ChatCompletionStreamResponse, ChatMessage,
|
|
ChatStreamChoice, ChatStreamDelta, ContentPart, ContentPartValue, FromOpenAI, ImageUrl, MessageContent,
|
|
OpenAIContentPart, OpenAIMessage, OpenAIRequest, ToOpenAI, UnifiedMessage, UnifiedRequest, Usage,
|
|
};
|
|
pub use providers::{Provider, ProviderManager, ProviderResponse, ProviderStreamChunk};
|
|
pub use server::router;
|
|
pub use state::AppState;
|
|
|
|
/// Test utilities for integration testing
|
|
#[cfg(test)]
|
|
pub mod test_utils {
|
|
use std::sync::Arc;
|
|
|
|
use crate::{client::ClientManager, providers::ProviderManager, rate_limiting::RateLimitManager, state::AppState, utils::crypto, database::run_migrations};
|
|
use sqlx::sqlite::SqlitePool;
|
|
|
|
/// Create a test application state
|
|
pub async fn create_test_state() -> AppState {
|
|
// Create in-memory database
|
|
let pool = SqlitePool::connect("sqlite::memory:")
|
|
.await
|
|
.expect("Failed to create test database");
|
|
|
|
// Run migrations on the pool
|
|
run_migrations(&pool).await.expect("Failed to run migrations");
|
|
|
|
let rate_limit_manager = RateLimitManager::new(
|
|
crate::rate_limiting::RateLimiterConfig::default(),
|
|
crate::rate_limiting::CircuitBreakerConfig::default(),
|
|
);
|
|
|
|
let client_manager = Arc::new(ClientManager::new(pool.clone()));
|
|
|
|
// Create provider manager
|
|
let provider_manager = ProviderManager::new();
|
|
|
|
let model_registry = crate::models::registry::ModelRegistry {
|
|
providers: std::collections::HashMap::new(),
|
|
};
|
|
|
|
let (dashboard_tx, _) = tokio::sync::broadcast::channel::<serde_json::Value>(100);
|
|
|
|
let config = Arc::new(crate::config::AppConfig {
|
|
server: crate::config::ServerConfig {
|
|
port: 8080,
|
|
host: "127.0.0.1".to_string(),
|
|
auth_tokens: vec![],
|
|
},
|
|
database: crate::config::DatabaseConfig {
|
|
path: std::path::PathBuf::from(":memory:"),
|
|
max_connections: 5,
|
|
},
|
|
providers: crate::config::ProviderConfig {
|
|
openai: crate::config::OpenAIConfig {
|
|
api_key_env: "OPENAI_API_KEY".to_string(),
|
|
base_url: "".to_string(),
|
|
default_model: "".to_string(),
|
|
enabled: true,
|
|
},
|
|
gemini: crate::config::GeminiConfig {
|
|
api_key_env: "GEMINI_API_KEY".to_string(),
|
|
base_url: "".to_string(),
|
|
default_model: "".to_string(),
|
|
enabled: true,
|
|
},
|
|
deepseek: crate::config::DeepSeekConfig {
|
|
api_key_env: "DEEPSEEK_API_KEY".to_string(),
|
|
base_url: "".to_string(),
|
|
default_model: "".to_string(),
|
|
enabled: true,
|
|
},
|
|
grok: crate::config::GrokConfig {
|
|
api_key_env: "GROK_API_KEY".to_string(),
|
|
base_url: "".to_string(),
|
|
default_model: "".to_string(),
|
|
enabled: true,
|
|
},
|
|
ollama: crate::config::OllamaConfig {
|
|
base_url: "".to_string(),
|
|
enabled: true,
|
|
models: vec![],
|
|
},
|
|
},
|
|
model_mapping: crate::config::ModelMappingConfig { patterns: vec![] },
|
|
pricing: crate::config::PricingConfig {
|
|
openai: vec![],
|
|
gemini: vec![],
|
|
deepseek: vec![],
|
|
grok: vec![],
|
|
ollama: vec![],
|
|
},
|
|
config_path: None,
|
|
encryption_key: "000102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f".to_string(),
|
|
});
|
|
|
|
// Initialize encryption with the test key
|
|
crypto::init_with_key(&config.encryption_key).expect("failed to initialize crypto");
|
|
|
|
AppState::new(
|
|
config,
|
|
provider_manager,
|
|
pool,
|
|
rate_limit_manager,
|
|
model_registry,
|
|
vec![], // auth_tokens
|
|
)
|
|
}
|
|
|
|
/// Create a test HTTP client
|
|
pub fn create_test_client() -> reqwest::Client {
|
|
reqwest::Client::builder()
|
|
.timeout(std::time::Duration::from_secs(30))
|
|
.build()
|
|
.expect("Failed to create test HTTP client")
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod integration_tests {
|
|
use super::test_utils::*;
|
|
use crate::{
|
|
models::{ChatCompletionRequest, ChatMessage},
|
|
server::router,
|
|
utils::crypto,
|
|
};
|
|
use axum::{
|
|
body::Body,
|
|
http::{Request, StatusCode},
|
|
};
|
|
use mockito::Server;
|
|
use serde_json::json;
|
|
use sqlx::Row;
|
|
use tower::util::ServiceExt;
|
|
|
|
#[tokio::test]
|
|
async fn test_encrypted_provider_key_integration() {
|
|
// Step 1: Setup test database and state
|
|
let state = create_test_state().await;
|
|
let pool = state.db_pool.clone();
|
|
|
|
// Step 2: Insert provider with encrypted API key
|
|
let test_api_key = "test-openai-key-12345";
|
|
let encrypted_key = crypto::encrypt(test_api_key).expect("Failed to encrypt test key");
|
|
|
|
sqlx::query(
|
|
r#"
|
|
INSERT INTO provider_configs (id, display_name, enabled, base_url, api_key, api_key_encrypted, credit_balance, low_credit_threshold)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
|
"#,
|
|
)
|
|
.bind("openai")
|
|
.bind("OpenAI")
|
|
.bind(true)
|
|
.bind("http://localhost:1234") // Mock server URL
|
|
.bind(&encrypted_key)
|
|
.bind(true) // api_key_encrypted flag
|
|
.bind(100.0)
|
|
.bind(5.0)
|
|
.execute(&pool)
|
|
.await
|
|
.expect("Failed to update provider URL");
|
|
|
|
// Re-initialize provider with new URL
|
|
state
|
|
.provider_manager
|
|
.initialize_provider("openai", &state.config, &pool)
|
|
.await
|
|
.expect("Failed to re-initialize provider");
|
|
|
|
// Step 4: Mock OpenAI API server
|
|
let mut server = Server::new_async().await;
|
|
let mock = server
|
|
.mock("POST", "/chat/completions")
|
|
.match_header("authorization", format!("Bearer {}", test_api_key).as_str())
|
|
.with_status(200)
|
|
.with_header("content-type", "application/json")
|
|
.with_body(
|
|
json!({
|
|
"id": "chatcmpl-test",
|
|
"object": "chat.completion",
|
|
"created": 1234567890,
|
|
"model": "gpt-3.5-turbo",
|
|
"choices": [{
|
|
"index": 0,
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": "Hello, world!"
|
|
},
|
|
"finish_reason": "stop"
|
|
}],
|
|
"usage": {
|
|
"prompt_tokens": 10,
|
|
"completion_tokens": 5,
|
|
"total_tokens": 15
|
|
}
|
|
})
|
|
.to_string(),
|
|
)
|
|
.create_async()
|
|
.await;
|
|
|
|
// Update provider base URL to use mock server
|
|
sqlx::query("UPDATE provider_configs SET base_url = ? WHERE id = 'openai'")
|
|
.bind(&server.url())
|
|
.execute(&pool)
|
|
.await
|
|
.expect("Failed to update provider URL");
|
|
|
|
// Re-initialize provider with new URL
|
|
state
|
|
.provider_manager
|
|
.initialize_provider("openai", &state.config, &pool)
|
|
.await
|
|
.expect("Failed to re-initialize provider");
|
|
|
|
// Step 5: Create test router and make request
|
|
let app = router(state);
|
|
|
|
let request_body = ChatCompletionRequest {
|
|
model: "gpt-3.5-turbo".to_string(),
|
|
messages: vec![ChatMessage {
|
|
role: "user".to_string(),
|
|
content: crate::models::MessageContent::Text {
|
|
content: "Hello".to_string(),
|
|
},
|
|
reasoning_content: None,
|
|
tool_calls: None,
|
|
name: None,
|
|
tool_call_id: None,
|
|
}],
|
|
temperature: None,
|
|
top_p: None,
|
|
top_k: None,
|
|
n: None,
|
|
stop: None,
|
|
max_tokens: Some(100),
|
|
presence_penalty: None,
|
|
frequency_penalty: None,
|
|
stream: Some(false),
|
|
tools: None,
|
|
tool_choice: None,
|
|
};
|
|
|
|
let request = Request::builder()
|
|
.method("POST")
|
|
.uri("/v1/chat/completions")
|
|
.header("content-type", "application/json")
|
|
.header("authorization", "Bearer test-token")
|
|
.body(Body::from(serde_json::to_string(&request_body).unwrap()))
|
|
.unwrap();
|
|
|
|
// Step 6: Execute request through proxy
|
|
let response = app
|
|
.oneshot(request)
|
|
.await
|
|
.expect("Failed to execute request");
|
|
|
|
let status = response.status();
|
|
println!("Response status: {}", status);
|
|
|
|
if status != StatusCode::OK {
|
|
let body_bytes = axum::body::to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
|
let body_str = String::from_utf8(body_bytes.to_vec()).unwrap();
|
|
println!("Response body: {}", body_str);
|
|
panic!("Response status is not OK: {}", status);
|
|
}
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
|
|
// Verify the mock was called
|
|
mock.assert_async().await;
|
|
|
|
// Give the async logging task time to complete
|
|
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
|
|
|
// Step 7: Verify usage was logged in database
|
|
let log_row = sqlx::query("SELECT * FROM llm_requests WHERE client_id = 'client_test-tok' ORDER BY id DESC LIMIT 1")
|
|
.fetch_one(&pool)
|
|
.await
|
|
.expect("Request log not found");
|
|
|
|
assert_eq!(log_row.get::<String, _>("provider"), "openai");
|
|
assert_eq!(log_row.get::<String, _>("model"), "gpt-3.5-turbo");
|
|
assert_eq!(log_row.get::<i64, _>("prompt_tokens"), 10);
|
|
assert_eq!(log_row.get::<i64, _>("completion_tokens"), 5);
|
|
assert_eq!(log_row.get::<i64, _>("total_tokens"), 15);
|
|
assert_eq!(log_row.get::<String, _>("status"), "success");
|
|
|
|
// Verify client usage was updated
|
|
let client_row = sqlx::query("SELECT total_requests, total_tokens, total_cost FROM clients WHERE client_id = 'client_test-tok'")
|
|
.fetch_one(&pool)
|
|
.await
|
|
.expect("Client not found");
|
|
|
|
assert_eq!(client_row.get::<i64, _>("total_requests"), 1);
|
|
assert_eq!(client_row.get::<i64, _>("total_tokens"), 15);
|
|
}
|
|
}
|