From 86895e14f41d15af7e8c4a9752c040f34b277497 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 16 Feb 2026 06:17:55 -0500 Subject: [PATCH 001/111] server: initial attempt --- server/src/api/chat.rs | 57 ++++++++++++++----- server/src/db/models/chat.rs | 3 + server/src/provider/core.rs | 22 ++++++- server/src/provider/providers/openai.rs | 1 + .../src/provider/providers/openai/request.rs | 4 +- .../src/provider/providers/openai/response.rs | 32 ++++++++++- server/src/storage/local.rs | 31 +++++++++- server/src/stream/tests/mod.rs | 36 +++++++++--- server/src/stream/tinistream.rs | 2 +- server/src/stream/writer.rs | 42 ++++++++------ server/src/utils/generate_title.rs | 1 + 11 files changed, 187 insertions(+), 44 deletions(-) diff --git a/server/src/api/chat.rs b/server/src/api/chat.rs index 33e8332..100987a 100644 --- a/server/src/api/chat.rs +++ b/server/src/api/chat.rs @@ -12,7 +12,7 @@ use crate::{ auth::ChatRsUserId, db::{ models::*, - services::{ChatDbService, ProviderDbService, ToolDbService}, + services::{ChatDbService, FileDbService, ProviderDbService, ToolDbService}, DbConnection, DbPool, }, errors::ApiError, @@ -184,32 +184,61 @@ pub async fn send_chat_stream( let tinistream = tinistream.inner().to_owned(); let provider_id = input.provider_id.clone(); let provider_options = input.options.clone(); + let storage = storage.inner().to_owned(); tokio::spawn(async move { - let mut stream_writer = LlmStreamWriter::new(); - let (text, tool_calls, usage, errors, cancelled) = - stream_writer.process(stream, ws_writer, ws_reader).await; + let response = LlmStreamWriter::new() + .process(stream, ws_writer, ws_reader) + .await; + + let mut image_ids: Option> = None; + for image in response.images.unwrap_or_default() { + let path = format!("generated/{}.png", Uuid::new_v4()); + match storage + .create_file_from_data_url(&user_id, Some(&session_id), &path, image.base64_url) + .await + { + Ok((content_type, size)) => { + match FileDbService::new(&mut db) + .create_session_file(NewChatRsFile { + user_id: &user_id, + session_id: Some(&session_id), + path: &path, + file_type: ChatRsFileType::Image.into(), + content_type: &content_type, + size: size.try_into().unwrap_or_default(), + }) + .await + { + Ok(file) => image_ids.get_or_insert_default().push(file.id), + Err(err) => rocket::error!("Failed to save image to db: {err}"), + } + } + Err(err) => rocket::error!("Failed to save image to storage: {err}"), + } + } let assistant_meta = AssistantMeta { provider_id, provider_options: Some(provider_options), - tool_calls, - usage, - errors, - partial: cancelled.then_some(true), + tool_calls: response.tool_calls, + images: image_ids, + usage: response.usage, + errors: response.errors, + partial: response.cancelled.then_some(true), }; - let db_result = ChatDbService::new(&mut db) + if let Err(err) = ChatDbService::new(&mut db) .save_message(NewChatRsMessage { session_id: &session_id, role: ChatRsMessageRole::Assistant, - content: &text.unwrap_or_default(), + content: &response.text.unwrap_or_default(), meta: ChatRsMessageMeta::new_assistant(assistant_meta), }) - .await; - if let Err(err) = db_result { - rocket::error!("Failed to save assistant message: {}", err); + .await + { + rocket::error!("Failed to save assistant message: {err}"); } - if !cancelled { + if !response.cancelled { tinistream.stream_end(&stream_key).await.ok(); } }); diff --git a/server/src/db/models/chat.rs b/server/src/db/models/chat.rs index 1b27574..0ef5a50 100644 --- a/server/src/db/models/chat.rs +++ b/server/src/db/models/chat.rs @@ -116,6 +116,9 @@ pub struct AssistantMeta { /// The tool calls requested by the assistant #[serde(skip_serializing_if = "Option::is_none")] pub tool_calls: Option>, + /// IDs of generated images + #[serde(skip_serializing_if = "Option::is_none")] + pub images: Option>, /// Provider usage information #[serde(skip_serializing_if = "Option::is_none")] pub usage: Option, diff --git a/server/src/provider/core.rs b/server/src/provider/core.rs index afbec27..332bda9 100644 --- a/server/src/provider/core.rs +++ b/server/src/provider/core.rs @@ -9,7 +9,7 @@ use uuid::Uuid; use crate::{ db::models::{ChatRsFileType, ChatRsToolCall}, - provider::models::LlmModel, + provider::models::{LlmModel, ModalityType}, }; /// Unified API for LLM providers @@ -42,6 +42,7 @@ pub enum LlmStreamChunk { Text(String), ToolCalls(Vec), PendingToolCall(LlmPendingToolCall), + Images(Vec), Usage(LlmUsage), } @@ -103,6 +104,12 @@ pub struct LlmPendingToolCall { pub tool_name: String, } +/// A generated image from the LLM provider +#[derive(Debug, Clone)] +pub struct LlmImage { + pub base64_url: String, +} + /// Usage stats from the LLM provider #[derive(Debug, Default, JsonSchema, serde::Serialize, serde::Deserialize)] pub struct LlmUsage { @@ -113,12 +120,25 @@ pub struct LlmUsage { pub cost: Option, } +/// Complete processed response from the LLM provider +pub struct LlmOutput { + pub text: Option, + pub tool_calls: Option>, + pub images: Option>, + pub usage: Option, + pub errors: Option>, + pub cancelled: bool, +} + /// Configuration for LLM provider requests #[derive(Clone, Debug, Default, JsonSchema, serde::Serialize, serde::Deserialize)] pub struct LlmProviderOptions { pub model: String, pub temperature: Option, pub max_tokens: Option, + /// Only supported for OpenRouter + #[serde(skip_serializing_if = "Option::is_none")] + pub modalities: Option>, } /// Generic message type to send to LLM providers diff --git a/server/src/provider/providers/openai.rs b/server/src/provider/providers/openai.rs index edf2c06..ac267aa 100644 --- a/server/src/provider/providers/openai.rs +++ b/server/src/provider/providers/openai.rs @@ -63,6 +63,7 @@ impl LlmApiProvider for OpenAIProvider { include_usage: true, }), tools: openai_tools, + modalities: options.modalities.clone(), }; let response = self diff --git a/server/src/provider/providers/openai/request.rs b/server/src/provider/providers/openai/request.rs index 87457d3..3b7f7a8 100644 --- a/server/src/provider/providers/openai/request.rs +++ b/server/src/provider/providers/openai/request.rs @@ -2,7 +2,7 @@ use serde::Serialize; use crate::{ db::models::ChatRsFileType, - provider::{utils::create_data_uri, LlmMessage, LlmTool}, + provider::{models::ModalityType, utils::create_data_uri, LlmMessage, LlmTool}, }; pub fn build_openai_messages<'a>(messages: &'a [LlmMessage]) -> Vec> { @@ -115,6 +115,8 @@ pub struct OpenAIRequest<'a> { pub stream_options: Option, #[serde(skip_serializing_if = "Option::is_none")] pub tools: Option>>, + #[serde(skip_serializing_if = "Option::is_none")] + pub modalities: Option>, } /// OpenAI API request stream options diff --git a/server/src/provider/providers/openai/response.rs b/server/src/provider/providers/openai/response.rs index 534d45d..667a169 100644 --- a/server/src/provider/providers/openai/response.rs +++ b/server/src/provider/providers/openai/response.rs @@ -2,7 +2,9 @@ use serde::Deserialize; use crate::{ db::models::ChatRsToolCall, - provider::{LlmPendingToolCall, LlmStreamChunk, LlmStreamChunkResult, LlmTool, LlmUsage}, + provider::{ + LlmImage, LlmPendingToolCall, LlmStreamChunk, LlmStreamChunkResult, LlmTool, LlmUsage, + }, }; /// Parse chunks from an OpenAI SSE event @@ -43,6 +45,16 @@ pub fn parse_openai_event( } } } + if let Some(images) = delta.images { + chunks.push(Ok(LlmStreamChunk::Images( + images + .into_iter() + .map(|image| LlmImage { + base64_url: image.image_url.url, + }) + .collect(), + ))); + } } if let Some(usage) = event.usage { chunks.push(Ok(LlmStreamChunk::Usage(usage.into()))); @@ -86,6 +98,9 @@ pub struct OpenAIResponseDelta { // role: Option, content: Option, tool_calls: Option>, + /// OpenRouter images + #[serde(skip_serializing_if = "Option::is_none")] + pub images: Option>, } /// OpenAI streaming tool call @@ -122,6 +137,21 @@ struct OpenAIStreamToolCallFunction { arguments: Option, } +/// OpenRouter image +#[derive(Debug, Deserialize)] +pub struct OpenRouterImage { + // #[serde(rename = "type")] + // pub image_type: String, + pub image_url: OpenRouterImageData, +} + +/// OpenRouter image data +#[derive(Debug, Deserialize)] +pub struct OpenRouterImageData { + /// Base64 data URL + pub url: String, +} + /// OpenAI API response usage #[derive(Debug, Deserialize)] pub struct OpenAIUsage { diff --git a/server/src/storage/local.rs b/server/src/storage/local.rs index 4c0cd9f..e83f577 100644 --- a/server/src/storage/local.rs +++ b/server/src/storage/local.rs @@ -8,6 +8,7 @@ use tokio::{ }; use uuid::Uuid; +#[derive(Debug, Clone)] pub struct LocalStorage { base_path: PathBuf, } @@ -72,6 +73,20 @@ impl LocalStorage { Ok(total_bytes_written) } + pub async fn create_file_from_data_url( + &self, + user_id: &Uuid, + session_id: Option<&Uuid>, + path: &str, + data_url: String, + ) -> IoResult<(String, u64)> { + let file_path = self.get_file_path(user_id, session_id, path)?; + let dir = file_path.parent().expect("Should have a parent directory"); + tokio::fs::create_dir_all(&dir).await?; + + tokio::task::spawn_blocking(move || save_base64_url(&data_url, &file_path)).await? + } + pub async fn delete_file>( &self, user_id: &Uuid, @@ -114,7 +129,6 @@ impl LocalStorage { } /// Synchronously read a file as a base64 encoded string. -/// (This is synchronous because the `base64` crate is synchronous.) fn read_base64(path: &Path) -> IoResult { let mut file = std::fs::File::open(path)?; let file_size = file.metadata()?.len(); @@ -132,3 +146,18 @@ fn read_base64(path: &Path) -> IoResult { } Ok(String::from_utf8(result).expect("base64 is valid UTF8")) } + +/// Synchronously save a base64 data URL to a file. Returns the content type and size of the saved file. +fn save_base64_url(data_url: &str, output_path: &Path) -> IoResult<(String, u64)> { + let (content_type, base64_data) = data_url + .split_once(',') + .ok_or(std::io::Error::other("Invalid data URL format"))?; + let mut decoder = base64::read::DecoderReader::new( + std::io::Cursor::new(base64_data.as_bytes()), + &base64::engine::general_purpose::STANDARD, + ); + let mut writer = std::io::BufWriter::new(std::fs::File::create(output_path)?); + let size = std::io::copy(&mut decoder, &mut writer)?; + + Ok((content_type.to_owned(), size)) +} diff --git a/server/src/stream/tests/mod.rs b/server/src/stream/tests/mod.rs index 41948e7..817b13a 100644 --- a/server/src/stream/tests/mod.rs +++ b/server/src/stream/tests/mod.rs @@ -7,8 +7,8 @@ use uuid::Uuid; use crate::{ provider::{ - providers::LoremProvider, LlmApiProvider, LlmProviderOptions, LlmStream, LlmStreamChunk, - LlmStreamError, LlmUsage, + providers::LoremProvider, LlmApiProvider, LlmOutput, LlmProviderOptions, LlmStream, + LlmStreamChunk, LlmStreamError, LlmUsage, }, stream::chat_stream_key, }; @@ -49,8 +49,14 @@ async fn stream_writer_basic_functionality() { .expect("Failed to create lorem stream"); // Process the stream - let (text, tool_calls, usage, errors, cancelled) = - writer.process(stream, ws_writer, ws_reader).await; + let LlmOutput { + text, + tool_calls, + usage, + errors, + cancelled, + .. + } = writer.process(stream, ws_writer, ws_reader).await; // Verify results assert!(text.is_some()); @@ -90,7 +96,9 @@ async fn stream_writer_batching() { ); let stream: LlmStream = Box::pin(chunk_stream); - let (text, _, _, _, cancelled) = writer.process(stream, ws_writer, ws_reader).await; + let LlmOutput { + text, cancelled, .. + } = writer.process(stream, ws_writer, ws_reader).await; assert!(text.is_some()); let text = text.unwrap(); @@ -116,7 +124,12 @@ async fn stream_writer_error_handling() { ]); let stream: LlmStream = Box::pin(error_stream); - let (text, _, _, errors, cancelled) = writer.process(stream, ws_writer, ws_reader).await; + let LlmOutput { + text, + errors, + cancelled, + .. + } = writer.process(stream, ws_writer, ws_reader).await; assert!(text.is_some()); let text = text.unwrap(); @@ -153,7 +166,9 @@ async fn stream_writer_cancel() { tini.stream_cancel(&key).await.unwrap(); // process() response should show that stream was cancelled - let (_, _, _, errors, cancelled) = process_fut.await; + let LlmOutput { + errors, cancelled, .. + } = process_fut.await; assert!(cancelled); assert!(errors.unwrap().last().unwrap().contains("cancelled")); @@ -188,7 +203,12 @@ async fn stream_writer_usage_tracking() { ]); let stream: LlmStream = Box::pin(usage_stream); - let (text, _, usage, _, cancelled) = writer.process(stream, ws_writer, ws_reader).await; + let LlmOutput { + text, + usage, + cancelled, + .. + } = writer.process(stream, ws_writer, ws_reader).await; assert!(text.is_some()); assert_eq!(text.unwrap(), "Hello World"); diff --git a/server/src/stream/tinistream.rs b/server/src/stream/tinistream.rs index 79700b9..e362fa1 100644 --- a/server/src/stream/tinistream.rs +++ b/server/src/stream/tinistream.rs @@ -115,7 +115,7 @@ impl TinistreamClient { Ok(res.into_inner().status) } - /// End a stream + /// Signal the end of a stream pub async fn stream_end(&self, key: &str) -> TiniResult { let res = self .client diff --git a/server/src/stream/writer.rs b/server/src/stream/writer.rs index 50fa097..7845906 100644 --- a/server/src/stream/writer.rs +++ b/server/src/stream/writer.rs @@ -10,7 +10,10 @@ use tokio_util::sync::CancellationToken; use crate::{ db::models::ChatRsToolCall, - provider::{LlmPendingToolCall, LlmStream, LlmStreamChunk, LlmStreamError, LlmUsage}, + provider::{ + LlmImage, LlmOutput, LlmPendingToolCall, LlmStream, LlmStreamChunk, LlmStreamError, + LlmUsage, + }, }; /// Interval at which chunks are flushed to the Redis stream. @@ -27,6 +30,8 @@ pub struct LlmStreamWriter { complete_text: Option, /// Accumulated tool calls from the assistant. tool_calls: Option>, + /// Accumulated generated images from the assistant. + images: Option>, /// Accumulated errors during the stream from the LLM provider. errors: Option>, /// Accumulated usage information from the LLM provider. @@ -58,6 +63,7 @@ impl LlmStreamWriter { current_chunk: ChunkState::default(), complete_text: None, tool_calls: None, + images: None, errors: None, usage: None, } @@ -71,13 +77,7 @@ impl LlmStreamWriter { stream: LlmStream, mut ws_writer: SplitSink, mut ws_reader: SplitStream, - ) -> ( - Option, - Option>, - Option, - Option>, - bool, - ) { + ) -> LlmOutput { let mut cancelled = false; // Spawn task to listen for stream cancellation @@ -102,15 +102,18 @@ impl LlmStreamWriter { cancel_task.abort(); ws_writer.close().await.ok(); - let complete_text = self.complete_text.take(); - let tool_calls = self.tool_calls.take(); - let usage = self.usage.take(); - let errors = self.errors.take().map(|e| { - e.into_iter() - .map(|e| e.to_string()) - .collect::>() - }); - (complete_text, tool_calls, usage, errors, cancelled) + LlmOutput { + text: self.complete_text.take(), + tool_calls: self.tool_calls.take(), + images: self.images.take(), + usage: self.usage.take(), + errors: self.errors.take().map(|e| { + e.into_iter() + .map(|e| e.to_string()) + .collect::>() + }), + cancelled, + } } async fn process_stream( @@ -127,6 +130,7 @@ impl LlmStreamWriter { LlmStreamChunk::PendingToolCall(pending_tool_call) => { self.process_pending_tool_call(pending_tool_call) } + LlmStreamChunk::Images(images) => self.process_images(images), LlmStreamChunk::Usage(usage) => self.process_usage(usage), }, Some(Err(err)) => self.process_error(err), @@ -174,6 +178,10 @@ impl LlmStreamWriter { } } + fn process_images(&mut self, images: Vec) { + self.images.get_or_insert_default().extend(images); + } + fn process_usage(&mut self, usage_chunk: LlmUsage) { let usage = self.usage.get_or_insert_default(); if let Some(input_tokens) = usage_chunk.input_tokens { diff --git a/server/src/utils/generate_title.rs b/server/src/utils/generate_title.rs index 2146608..0108780 100644 --- a/server/src/utils/generate_title.rs +++ b/server/src/utils/generate_title.rs @@ -47,6 +47,7 @@ async fn generate( model, temperature: Some(DEFAULT_TEMPERATURE), max_tokens: Some(TITLE_TOKENS), + ..Default::default() }; let message = format!("{}: \"{}\"", TITLE_PROMPT, user_message); let title = provider.prompt(&message, &provider_options).await?; From 3c2da7de9457222e501497ebc02b770b1749e11b Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 16 Feb 2026 07:10:33 -0500 Subject: [PATCH 002/111] server: refactor stream and persist logic out of API controller --- server/src/api/chat.rs | 83 +++++-------------------- server/src/storage/local.rs | 1 + server/src/stream.rs | 2 + server/src/stream/streamer.rs | 114 ++++++++++++++++++++++++++++++++++ server/src/stream/writer.rs | 4 +- 5 files changed, 133 insertions(+), 71 deletions(-) create mode 100644 server/src/stream/streamer.rs diff --git a/server/src/api/chat.rs b/server/src/api/chat.rs index 100987a..6a78fa8 100644 --- a/server/src/api/chat.rs +++ b/server/src/api/chat.rs @@ -1,6 +1,6 @@ use std::borrow::Cow; -use rocket::{futures::StreamExt, get, post, serde::json::Json, Route, State}; +use rocket::{get, post, serde::json::Json, Route, State}; use rocket_okapi::{ okapi::openapi3::OpenApi, openapi, openapi_get_routes_spec, settings::OpenApiSettings, }; @@ -12,7 +12,7 @@ use crate::{ auth::ChatRsUserId, db::{ models::*, - services::{ChatDbService, FileDbService, ProviderDbService, ToolDbService}, + services::{ChatDbService, ProviderDbService, ToolDbService}, DbConnection, DbPool, }, errors::ApiError, @@ -169,79 +169,24 @@ pub async fn send_chat_stream( messages.push(message); } - // Convert the messages, and get the provider's response + // Convert the messages, and get the streaming response from the provider let llm_messages = build_llm_messages(messages, &user_id, &session_id, &mut db, &storage).await?; let stream = provider_api .chat_stream(llm_messages, tools, &input.options) .await?; - // Create the Redis stream and get a WebSocket connection for writing to it - let stream_access = tinistream.stream_start(&stream_key).await?; - let (ws_writer, ws_reader) = tinistream.stream_writer_ws(&stream_key).await?.split(); - - // Spawn a task to stream and save the response - let tinistream = tinistream.inner().to_owned(); - let provider_id = input.provider_id.clone(); - let provider_options = input.options.clone(); - let storage = storage.inner().to_owned(); - tokio::spawn(async move { - let response = LlmStreamWriter::new() - .process(stream, ws_writer, ws_reader) - .await; - - let mut image_ids: Option> = None; - for image in response.images.unwrap_or_default() { - let path = format!("generated/{}.png", Uuid::new_v4()); - match storage - .create_file_from_data_url(&user_id, Some(&session_id), &path, image.base64_url) - .await - { - Ok((content_type, size)) => { - match FileDbService::new(&mut db) - .create_session_file(NewChatRsFile { - user_id: &user_id, - session_id: Some(&session_id), - path: &path, - file_type: ChatRsFileType::Image.into(), - content_type: &content_type, - size: size.try_into().unwrap_or_default(), - }) - .await - { - Ok(file) => image_ids.get_or_insert_default().push(file.id), - Err(err) => rocket::error!("Failed to save image to db: {err}"), - } - } - Err(err) => rocket::error!("Failed to save image to storage: {err}"), - } - } - - let assistant_meta = AssistantMeta { - provider_id, - provider_options: Some(provider_options), - tool_calls: response.tool_calls, - images: image_ids, - usage: response.usage, - errors: response.errors, - partial: response.cancelled.then_some(true), - }; - if let Err(err) = ChatDbService::new(&mut db) - .save_message(NewChatRsMessage { - session_id: &session_id, - role: ChatRsMessageRole::Assistant, - content: &response.text.unwrap_or_default(), - meta: ChatRsMessageMeta::new_assistant(assistant_meta), - }) - .await - { - rocket::error!("Failed to save assistant message: {err}"); - } - - if !response.cancelled { - tinistream.stream_end(&stream_key).await.ok(); - } - }); + // Start the client stream and get the access URL / token + let stream_access = LlmClientStreamer::new(db, tinistream, storage) + .start( + stream, + stream_key, + user_id.clone(), + session_id, + input.provider_id, + input.into_inner().options, + ) + .await?; Ok(Json(StreamAccess { url: stream_access.sse_url, diff --git a/server/src/storage/local.rs b/server/src/storage/local.rs index e83f577..806d173 100644 --- a/server/src/storage/local.rs +++ b/server/src/storage/local.rs @@ -8,6 +8,7 @@ use tokio::{ }; use uuid::Uuid; +/// Local file storage #[derive(Debug, Clone)] pub struct LocalStorage { base_path: PathBuf, diff --git a/server/src/stream.rs b/server/src/stream.rs index dc6ed00..0c31b3c 100644 --- a/server/src/stream.rs +++ b/server/src/stream.rs @@ -1,9 +1,11 @@ #[cfg(test)] mod tests; +mod streamer; mod tinistream; mod writer; +pub use streamer::*; pub use tinistream::*; pub use writer::*; diff --git a/server/src/stream/streamer.rs b/server/src/stream/streamer.rs new file mode 100644 index 0000000..f49b1ac --- /dev/null +++ b/server/src/stream/streamer.rs @@ -0,0 +1,114 @@ +use rocket::futures::StreamExt; +use tinistream_client::types::StreamAccessResponse; +use uuid::Uuid; + +use crate::{ + db::{ + models::{ + AssistantMeta, ChatRsFileType, ChatRsMessageMeta, ChatRsMessageRole, NewChatRsFile, + NewChatRsMessage, + }, + services::{ChatDbService, FileDbService}, + DbConnection, + }, + errors::ApiError, + provider::{LlmProviderOptions, LlmStream}, + storage::LocalStorage, + stream::TinistreamClient, +}; + +/// Utility that handles streaming to clients and persisting responses from the provider +pub struct LlmClientStreamer { + db: DbConnection, + tinistream: TinistreamClient, + storage: LocalStorage, +} + +impl LlmClientStreamer { + pub fn new(db: DbConnection, tinistream: &TinistreamClient, storage: &LocalStorage) -> Self { + Self { + db, + tinistream: tinistream.to_owned(), + storage: storage.to_owned(), + } + } + + pub async fn start( + mut self, + stream: LlmStream, + stream_key: String, + user_id: Uuid, + session_id: Uuid, + provider_id: i32, + provider_options: LlmProviderOptions, + ) -> Result { + // Create the Redis stream in `tinistream` and get a WebSocket connection for writing to it + let stream_access = self.tinistream.stream_start(&stream_key).await?; + let (ws_writer, ws_reader) = self.tinistream.stream_writer_ws(&stream_key).await?.split(); + + // Spawn a task to finish streaming and process/save the response + tokio::spawn(async move { + let response = super::LlmStreamWriter::new() + .process(stream, ws_writer, ws_reader) + .await; + + // Save generated images + let mut image_ids: Option> = None; + for image in response.images.unwrap_or_default() { + let path = format!("generated/{}.png", Uuid::new_v4()); + match self + .storage + .create_file_from_data_url(&user_id, Some(&session_id), &path, image.base64_url) + .await + { + Ok((content_type, size)) => { + match FileDbService::new(&mut self.db) + .create_session_file(NewChatRsFile { + user_id: &user_id, + session_id: Some(&session_id), + path: &path, + file_type: ChatRsFileType::Image.into(), + content_type: &content_type, + size: size.try_into().unwrap_or_default(), + }) + .await + { + Ok(file) => image_ids.get_or_insert_default().push(file.id), + Err(err) => rocket::error!("Failed to save image to db: {err}"), + } + } + Err(err) => rocket::error!("Failed to save image to storage: {err}"), + } + } + + // Save response message and metadata + let assistant_meta = AssistantMeta { + provider_id, + provider_options: Some(provider_options), + tool_calls: response.tool_calls, + images: image_ids, + usage: response.usage, + errors: response.errors, + partial: response.cancelled.then_some(true), + }; + if let Err(err) = ChatDbService::new(&mut self.db) + .save_message(NewChatRsMessage { + session_id: &session_id, + role: ChatRsMessageRole::Assistant, + content: &response.text.unwrap_or_default(), + meta: ChatRsMessageMeta::new_assistant(assistant_meta), + }) + .await + { + rocket::error!("Failed to save assistant message: {err}"); + } + + // Signal end of stream + if !response.cancelled { + self.tinistream.stream_end(&stream_key).await.ok(); + } + }); + + Ok(stream_access) + } +} diff --git a/server/src/stream/writer.rs b/server/src/stream/writer.rs index 7845906..8fb0da5 100644 --- a/server/src/stream/writer.rs +++ b/server/src/stream/writer.rs @@ -21,7 +21,7 @@ const FLUSH_INTERVAL: Duration = Duration::from_millis(400); /// Max # of characters of the text chunk before it is automatically flushed to Redis. const MAX_CHUNK_SIZE: usize = 200; -/// Utility for processing an incoming LLM response stream and writing to a Redis stream. +/// Utility for processing an incoming LLM response stream and writing chunks to `tinistream`. #[derive(Debug)] pub struct LlmStreamWriter { /// The current chunk of data being processed. @@ -70,7 +70,7 @@ impl LlmStreamWriter { } /// Process the incoming stream from the LLM provider, intermittently flushing - /// chunks to tinistream via the WebSocket connection, and return the final + /// chunks to `tinistream` via the WebSocket connection, and return the final /// accumulated response. pub async fn process( &mut self, From 08b0f266fb80ab2f55478ce0c22c4db33a805614 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 16 Feb 2026 07:17:03 -0500 Subject: [PATCH 003/111] Update chat.rs --- server/src/api/chat.rs | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/server/src/api/chat.rs b/server/src/api/chat.rs index 6a78fa8..645e094 100644 --- a/server/src/api/chat.rs +++ b/server/src/api/chat.rs @@ -169,7 +169,7 @@ pub async fn send_chat_stream( messages.push(message); } - // Convert the messages, and get the streaming response from the provider + // Build the messages and get the initial stream response from the provider let llm_messages = build_llm_messages(messages, &user_id, &session_id, &mut db, &storage).await?; let stream = provider_api @@ -204,11 +204,11 @@ pub async fn connect_to_chat_stream( tinistream: &State, ) -> Result, ApiError> { let key = chat_stream_key(&user_id, &session_id); - let connect = tinistream.stream_connect(&key).await?; + let stream_access = tinistream.stream_connect(&key).await?; Ok(Json(StreamAccess { - url: connect.sse_url, - token: connect.token, + url: stream_access.sse_url, + token: stream_access.token, })) } From 36fae7a96ff01821b7c34a2ef764c62a34288100 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 16 Feb 2026 07:24:33 -0500 Subject: [PATCH 004/111] tweaks --- server/src/provider/providers/openai.rs | 15 +++++++++------ server/src/provider/providers/openai/request.rs | 2 +- 2 files changed, 10 insertions(+), 7 deletions(-) diff --git a/server/src/provider/providers/openai.rs b/server/src/provider/providers/openai.rs index ac267aa..cb7373b 100644 --- a/server/src/provider/providers/openai.rs +++ b/server/src/provider/providers/openai.rs @@ -50,12 +50,15 @@ impl LlmApiProvider for OpenAIProvider { let request = OpenAIRequest { model: &options.model, messages: openai_messages, - max_tokens: (options.max_tokens.is_some() && self.base_url != OPENAI_API_BASE_URL) - .then(|| options.max_tokens.expect("already checked for Some value")), // OpenAI official API has deprecated `max_tokens` for `max_completion_tokens` - max_completion_tokens: (options.max_tokens.is_some() - && self.base_url == OPENAI_API_BASE_URL) - .then(|| options.max_tokens.expect("already checked for Some value")), + max_tokens: match options.max_tokens { + Some(max_tokens) if self.base_url != OPENAI_API_BASE_URL => Some(max_tokens), + _ => None, + }, + max_completion_tokens: match options.max_tokens { + Some(max_tokens) if self.base_url == OPENAI_API_BASE_URL => Some(max_tokens), + _ => None, + }, temperature: options.temperature, store: (self.base_url == OPENAI_API_BASE_URL).then_some(false), stream: Some(true), @@ -63,7 +66,7 @@ impl LlmApiProvider for OpenAIProvider { include_usage: true, }), tools: openai_tools, - modalities: options.modalities.clone(), + modalities: options.modalities.as_ref(), }; let response = self diff --git a/server/src/provider/providers/openai/request.rs b/server/src/provider/providers/openai/request.rs index 3b7f7a8..6e2b32d 100644 --- a/server/src/provider/providers/openai/request.rs +++ b/server/src/provider/providers/openai/request.rs @@ -116,7 +116,7 @@ pub struct OpenAIRequest<'a> { #[serde(skip_serializing_if = "Option::is_none")] pub tools: Option>>, #[serde(skip_serializing_if = "Option::is_none")] - pub modalities: Option>, + pub modalities: Option<&'a Vec>, } /// OpenAI API request stream options From f6c279de251b10e4f5bc83e57b83e3596966edf7 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 16 Feb 2026 13:20:55 -0500 Subject: [PATCH 005/111] rename --- server/src/db/models/chat.rs | 4 ++-- server/src/stream/streamer.rs | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/server/src/db/models/chat.rs b/server/src/db/models/chat.rs index 0ef5a50..371c21f 100644 --- a/server/src/db/models/chat.rs +++ b/server/src/db/models/chat.rs @@ -116,9 +116,9 @@ pub struct AssistantMeta { /// The tool calls requested by the assistant #[serde(skip_serializing_if = "Option::is_none")] pub tool_calls: Option>, - /// IDs of generated images + /// IDs of generated files #[serde(skip_serializing_if = "Option::is_none")] - pub images: Option>, + pub files: Option>, /// Provider usage information #[serde(skip_serializing_if = "Option::is_none")] pub usage: Option, diff --git a/server/src/stream/streamer.rs b/server/src/stream/streamer.rs index f49b1ac..49981b8 100644 --- a/server/src/stream/streamer.rs +++ b/server/src/stream/streamer.rs @@ -86,7 +86,7 @@ impl LlmClientStreamer { provider_id, provider_options: Some(provider_options), tool_calls: response.tool_calls, - images: image_ids, + files: image_ids, usage: response.usage, errors: response.errors, partial: response.cancelled.then_some(true), From 012c3c066062f81540d641fc9c2e584410d94afb Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 16 Feb 2026 13:37:13 -0500 Subject: [PATCH 006/111] perf: use pipelined database queries where possible new version of diesel_async allows for pipelined queries --- server/src/db/services/chat.rs | 25 ++++++++++++++----------- server/src/db/services/tool.rs | 22 ++++++++++++---------- 2 files changed, 26 insertions(+), 21 deletions(-) diff --git a/server/src/db/services/chat.rs b/server/src/db/services/chat.rs index 5fdb0b8..360373d 100644 --- a/server/src/db/services/chat.rs +++ b/server/src/db/services/chat.rs @@ -1,5 +1,6 @@ use diesel::prelude::*; use diesel_async::RunQueryDsl; +use rocket::futures; use uuid::Uuid; use crate::{ @@ -121,17 +122,19 @@ impl<'a> ChatDbService<'a> { user_id: &Uuid, session_id: &Uuid, ) -> Result<(ChatRsSession, Vec), diesel::result::Error> { - let session = chat_sessions::table - .filter(chat_sessions::user_id.eq(user_id)) - .filter(chat_sessions::id.eq(session_id)) - .select(ChatRsSession::as_select()) - .first(self.db) - .await?; - let messages = ChatRsMessage::belonging_to(&session) - .select(ChatRsMessage::as_select()) - .order_by(chat_messages::created_at.asc()) - .load(self.db) - .await?; + let (session, messages) = futures::future::try_join( + chat_sessions::table + .filter(chat_sessions::user_id.eq(user_id)) + .filter(chat_sessions::id.eq(session_id)) + .select(ChatRsSession::as_select()) + .first(self.db), + chat_messages::table + .filter(chat_messages::session_id.eq(session_id)) + .select(ChatRsMessage::as_select()) + .order_by(chat_messages::created_at.asc()) + .load(self.db), + ) + .await?; Ok((session, messages)) } diff --git a/server/src/db/services/tool.rs b/server/src/db/services/tool.rs index 98eee3e..4533eb5 100644 --- a/server/src/db/services/tool.rs +++ b/server/src/db/services/tool.rs @@ -1,6 +1,7 @@ use diesel::prelude::*; use diesel::result::Error; use diesel_async::RunQueryDsl; +use rocket::futures; use uuid::Uuid; use crate::db::{ @@ -25,16 +26,17 @@ impl<'a> ToolDbService<'a> { &mut self, user_id: &Uuid, ) -> Result<(Vec, Vec), Error> { - let system_tools = system_tools::table - .filter(system_tools::user_id.eq(user_id)) - .select(ChatRsSystemTool::as_select()) - .load(self.db) - .await?; - let external_api_tools = external_api_tools::table - .filter(external_api_tools::user_id.eq(user_id)) - .select(ChatRsExternalApiTool::as_select()) - .load(self.db) - .await?; + let (system_tools, external_api_tools) = futures::future::try_join( + system_tools::table + .filter(system_tools::user_id.eq(user_id)) + .select(ChatRsSystemTool::as_select()) + .load(self.db), + external_api_tools::table + .filter(external_api_tools::user_id.eq(user_id)) + .select(ChatRsExternalApiTool::as_select()) + .load(self.db), + ) + .await?; Ok((system_tools, external_api_tools)) } From e4f37c6cc38bda80d6a5684dc50e0f10e1510185 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 16 Feb 2026 13:39:23 -0500 Subject: [PATCH 007/111] Update db.rs --- server/src/db.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/server/src/db.rs b/server/src/db.rs index 32f9e14..5340788 100644 --- a/server/src/db.rs +++ b/server/src/db.rs @@ -55,7 +55,7 @@ impl<'r> FromRequest<'r> for DbConnection { Ok(conn) => Outcome::Success(DbConnection(conn)), Err(e) => { rocket::error!("Couldn't get database connection: {e}"); - Outcome::Error((Status::InternalServerError, "Couldn't get connection")) + Outcome::Error((Status::InternalServerError, "Couldn't get db connection")) } } } From 067c02e0e6667b99e444e0a5eac76dafe26f7778 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 16 Feb 2026 13:46:28 -0500 Subject: [PATCH 008/111] Update db.rs --- server/src/db.rs | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/server/src/db.rs b/server/src/db.rs index 5340788..4bc7089 100644 --- a/server/src/db.rs +++ b/server/src/db.rs @@ -64,24 +64,24 @@ impl<'r> FromRequest<'r> for DbConnection { /// Fairing that sets up and initializes the Postgres database pub fn setup_db() -> AdHoc { AdHoc::on_ignite("Database", |rocket| async { - let config = AsyncDieselConnectionManager::::new( - &get_app_config(&rocket).database_url, - ); + let db_url = get_app_config(&rocket).database_url.as_str(); + let config = AsyncDieselConnectionManager::::new(db_url); let pool: DbPool = Pool::builder(config) .build() .expect("Failed to parse database URL"); const MIGRATIONS: EmbeddedMigrations = embed_migrations!(); - let cxn = pool.get().await.expect("Failed to connect to database"); - tokio::task::spawn_blocking(move || { - AsyncConnectionWrapper::>::from(cxn) + let migration_cxn = pool.get().await.expect("Failed to connect to database"); + match tokio::task::spawn_blocking(move || { + AsyncConnectionWrapper::>::from(migration_cxn) .run_pending_migrations(MIGRATIONS) .expect("Database migrations failed"); }) .await - .expect("Database migration task failed"); - - rocket::info!("Migrations completed successfully"); + { + Ok(_) => rocket::info!("Migrations completed successfully"), + Err(err) => panic!("Database migration task failed: {err}"), + }; let shutdown = AdHoc::on_shutdown("Shutdown database", |rocket| { Box::pin(async { From 81230298192f937eb40cbef8e7a0104813514e05 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 16 Feb 2026 15:11:56 -0500 Subject: [PATCH 009/111] fix data url decoding and saving --- server/src/storage/local.rs | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/server/src/storage/local.rs b/server/src/storage/local.rs index 806d173..3913806 100644 --- a/server/src/storage/local.rs +++ b/server/src/storage/local.rs @@ -1,5 +1,5 @@ use std::{ - io::Result as IoResult, + io::{Result as IoResult, Write}, path::{Path, PathBuf}, }; use tokio::{ @@ -150,15 +150,21 @@ fn read_base64(path: &Path) -> IoResult { /// Synchronously save a base64 data URL to a file. Returns the content type and size of the saved file. fn save_base64_url(data_url: &str, output_path: &Path) -> IoResult<(String, u64)> { - let (content_type, base64_data) = data_url + let (prefix, base64_data) = data_url .split_once(',') .ok_or(std::io::Error::other("Invalid data URL format"))?; + let content_type = prefix + .strip_prefix("data:") + .and_then(|p| p.strip_suffix(";base64")) + .ok_or(std::io::Error::other("Invalid data URL prefix"))?; + let mut decoder = base64::read::DecoderReader::new( std::io::Cursor::new(base64_data.as_bytes()), &base64::engine::general_purpose::STANDARD, ); let mut writer = std::io::BufWriter::new(std::fs::File::create(output_path)?); - let size = std::io::copy(&mut decoder, &mut writer)?; + let total_bytes = std::io::copy(&mut decoder, &mut writer)?; + writer.flush()?; - Ok((content_type.to_owned(), size)) + Ok((content_type.to_owned(), total_bytes)) } From 76e725ad60babe1acd44cf01edda3d9175098d03 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 17 Feb 2026 01:44:04 -0500 Subject: [PATCH 010/111] web: enable image responses (more work needed) --- web/src/components/chat/ChatMessageInput.tsx | 62 +++++++++++-------- .../components/chat/messages/ChatMessage.tsx | 13 ++++ .../chat/settings/ChatModelSelect.tsx | 23 ++++--- .../chat/settings/ChatMoreSettings.tsx | 58 +++++++++-------- .../chat/settings/ChatProviderSelect.tsx | 2 +- .../chat/settings/ChatSettingsBadge.tsx | 2 +- web/src/hooks/useChatInputState.tsx | 29 +++++++-- web/src/lib/api/types.d.ts | 4 ++ 8 files changed, 123 insertions(+), 70 deletions(-) diff --git a/web/src/components/chat/ChatMessageInput.tsx b/web/src/components/chat/ChatMessageInput.tsx index 018665d..84bb4fe 100644 --- a/web/src/components/chat/ChatMessageInput.tsx +++ b/web/src/components/chat/ChatMessageInput.tsx @@ -1,7 +1,7 @@ import { CornerDownLeft, Paperclip, Upload, X } from "lucide-react"; import { - type FormEventHandler, memo, + type SubmitEventHandler, useCallback, useMemo, useState, @@ -35,9 +35,10 @@ export default memo(function ChatMessageInput({ const isMobile = useIsMobile(); const { + sessionId, providerId, modelId, - sessionId, + selectedModel, toolInput, files, maxTokens, @@ -84,7 +85,7 @@ export default memo(function ChatMessageInput({ [enterKeyShouldSubmit, onSubmitUserMessage], ); - const handleFormSubmit: FormEventHandler = useCallback( + const handleFormSubmit: SubmitEventHandler = useCallback( (ev) => { ev.preventDefault(); onSubmitUserMessage(); @@ -197,32 +198,39 @@ export default memo(function ChatMessageInput({ currentTemperature={temperature} onSelectMaxTokens={setMaxTokens} onSelectTemperature={setTemperature} + showTemperature={selectedModel?.temperature} /> - - - {files.length > 0 && ( -
- - {files.map((file) => file.path).join(", ")} -
+ {(!selectedModel || selectedModel.tool_call) && ( + )} - {uploadingFiles.length > 0 && ( -
- - Uploading... -
+ {(!selectedModel || selectedModel.attachment) && ( + <> + + {files.length > 0 && ( +
+ + {files.map((file) => file.path).join(", ")} +
+ )} + {uploadingFiles.length > 0 && ( +
+ + Uploading... +
+ )} + )} )} diff --git a/web/src/components/chat/messages/ChatMessage.tsx b/web/src/components/chat/messages/ChatMessage.tsx index 8162e0d..bb28b70 100644 --- a/web/src/components/chat/messages/ChatMessage.tsx +++ b/web/src/components/chat/messages/ChatMessage.tsx @@ -86,6 +86,19 @@ export default function ChatMessage({ onExecute={(id) => onExecuteToolCall(message.id, id)} /> )} + {message.meta.assistant?.files?.map((fileId) => ( + + ))}
diff --git a/web/src/components/chat/settings/ChatModelSelect.tsx b/web/src/components/chat/settings/ChatModelSelect.tsx index 7d08e7c..c1c5cb5 100644 --- a/web/src/components/chat/settings/ChatModelSelect.tsx +++ b/web/src/components/chat/settings/ChatModelSelect.tsx @@ -1,4 +1,4 @@ -import { Check, ChevronsUpDown } from "lucide-react"; +import { ChevronsUpDown, Eye, FileText, ImageIcon, Wrench } from "lucide-react"; import React from "react"; import PopoverDrawer from "@/components/PopoverDrawer"; @@ -37,7 +37,7 @@ export default function ChatModelSelect({ variant="outline" role="combobox" aria-expanded={open} - className="w-[180px] md:w-[240px] justify-between" + className="w-45 md:w-60 justify-between" > {currentModelId @@ -63,18 +63,23 @@ export default function ChatModelSelect({ setOpen(false); }} > -
+
{model.name} {model.id}
- +
+ {model.tool_call && } + {model.modalities?.input.includes("image") && } + {model.modalities?.input.includes("pdf") && } + {model.modalities?.output.includes("image") && } +
))} diff --git a/web/src/components/chat/settings/ChatMoreSettings.tsx b/web/src/components/chat/settings/ChatMoreSettings.tsx index 52694c9..6e1259e 100644 --- a/web/src/components/chat/settings/ChatMoreSettings.tsx +++ b/web/src/components/chat/settings/ChatMoreSettings.tsx @@ -18,6 +18,7 @@ interface Props { onSelectMaxTokens: (tokens: number) => void; currentTemperature: number; onSelectTemperature: (temperature: number) => void; + showTemperature?: boolean | null; } export default function ChatMoreSettings({ @@ -25,6 +26,7 @@ export default function ChatMoreSettings({ onSelectMaxTokens, currentTemperature, onSelectTemperature, + showTemperature = true, }: Props) { return ( onSelectMaxTokens(+tokens)} > - + @@ -57,32 +59,34 @@ export default function ChatMoreSettings({ - + {showTemperature && ( + + )}
); diff --git a/web/src/components/chat/settings/ChatProviderSelect.tsx b/web/src/components/chat/settings/ChatProviderSelect.tsx index a5bc878..677f076 100644 --- a/web/src/components/chat/settings/ChatProviderSelect.tsx +++ b/web/src/components/chat/settings/ChatProviderSelect.tsx @@ -37,7 +37,7 @@ export default function ChatProviderSelect({ variant="outline" role="combobox" aria-expanded={open} - className="w-[130px] md:w-[160px] justify-between" + className="w-32.5 md:w-40 justify-between" > {currentProvider ? currentProvider.name : "Select provider"} diff --git a/web/src/components/chat/settings/ChatSettingsBadge.tsx b/web/src/components/chat/settings/ChatSettingsBadge.tsx index fae864f..79d0d36 100644 --- a/web/src/components/chat/settings/ChatSettingsBadge.tsx +++ b/web/src/components/chat/settings/ChatSettingsBadge.tsx @@ -6,7 +6,7 @@ export default function ChatSettingsBadge({ children: React.ReactNode; }) { return ( - + {children} ); diff --git a/web/src/hooks/useChatInputState.tsx b/web/src/hooks/useChatInputState.tsx index cba1896..f777dff 100644 --- a/web/src/hooks/useChatInputState.tsx +++ b/web/src/hooks/useChatInputState.tsx @@ -1,5 +1,6 @@ import { useCallback, useEffect, useMemo, useRef, useState } from "react"; +import { useProviderModels } from "@/lib/api/provider"; import type { components } from "@/lib/api/types"; const DEFAULT_MAX_TOKENS = 2000; @@ -40,17 +41,26 @@ export const useChatInputState = ({ () => providers?.find((p) => p.id === providerId), [providers, providerId], ); + + const { data: models } = useProviderModels(providerId); const [modelId, setModel] = useState(initialOptions?.model || ""); + const selectedModel = useMemo( + () => models?.find((m) => m.id === modelId), + [models, modelId], + ); + const [toolInput, setToolInput] = useState< components["schemas"]["SendChatToolInput"] | null >(initialTools || DEFAULT_TOOL_INPUT); const [files, setFiles] = useState([]); + const [maxTokens, setMaxTokens] = useState( initialOptions?.max_tokens ?? DEFAULT_MAX_TOKENS, ); const [temperature, setTemperature] = useState( initialOptions?.temperature ?? DEFAULT_TEMPERATURE, ); + const [error, setError] = useState(""); // Reset state when session changes @@ -151,16 +161,21 @@ export const useChatInputState = ({ provider_id: providerId, options: { model: modelId, - temperature, + temperature: selectedModel?.temperature ? temperature : undefined, max_tokens: maxTokens, + modalities: selectedModel?.modalities?.output, }, - tools: toolInput, - files: files.length > 0 ? files.map((file) => file.id) : undefined, + tools: selectedModel?.tool_call ? toolInput : undefined, + files: + selectedModel?.attachment && files.length > 0 + ? files.map((file) => file.id) + : undefined, }); formRef.current?.reset(); }, [ providerId, selectedProvider, + selectedModel, modelId, toolInput, files, @@ -178,14 +193,16 @@ export const useChatInputState = ({ provider_id: providerId, options: { model: modelId, - temperature, + temperature: selectedModel?.temperature ? temperature : undefined, max_tokens: maxTokens, + modalities: selectedModel?.modalities?.output, }, - tools: toolInput, + tools: selectedModel?.tool_call ? toolInput : undefined, }); }, [ providerId, modelId, + selectedModel, toolInput, temperature, maxTokens, @@ -198,6 +215,7 @@ export const useChatInputState = ({ providerId, modelId, sessionId, + selectedModel, toolInput, files, maxTokens, @@ -222,6 +240,7 @@ export const useChatInputState = ({ providerId, modelId, sessionId, + selectedModel, toolInput, files, maxTokens, diff --git a/web/src/lib/api/types.d.ts b/web/src/lib/api/types.d.ts index 8cb7c4b..773a8cc 100644 --- a/web/src/lib/api/types.d.ts +++ b/web/src/lib/api/types.d.ts @@ -781,6 +781,8 @@ export interface components { provider_options?: components["schemas"]["LlmProviderOptions"] | null; /** @description The tool calls requested by the assistant */ tool_calls?: components["schemas"]["ChatRsToolCall"][] | null; + /** @description IDs of generated files */ + files?: string[] | null; /** @description Provider usage information */ usage?: components["schemas"]["LlmUsage"] | null; /** @description Errors encountered during message generation */ @@ -795,6 +797,8 @@ export interface components { temperature?: number | null; /** Format: uint32 */ max_tokens?: number | null; + /** @description Only supported for OpenRouter */ + modalities?: components["schemas"]["ModalityType"][] | null; }; /** @description A tool call requested by the provider */ ChatRsToolCall: { From 6d6f30ca378bae3cc160ffb56fb6a2a394d2c813 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 17 Feb 2026 01:46:19 -0500 Subject: [PATCH 011/111] server: switch back to using default model for session title generation --- server/src/api/chat.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/server/src/api/chat.rs b/server/src/api/chat.rs index 645e094..5722fe3 100644 --- a/server/src/api/chat.rs +++ b/server/src/api/chat.rs @@ -151,7 +151,7 @@ pub async fn send_chat_stream( &session_id, &user_message, &provider_api, - &input.options.model, + &provider.default_model, db_pool, ); } From 25d89fbca0e7a0c6d708b8adbd1bb1981218e056 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 17 Feb 2026 21:10:38 -0500 Subject: [PATCH 012/111] update readmes --- ARCHITECTURE.md | 19 ++++++------------ README.md | 48 ++++++++++++++++++++++++++++++--------------- server/.env.example | 7 ++++++- 3 files changed, 44 insertions(+), 30 deletions(-) diff --git a/ARCHITECTURE.md b/ARCHITECTURE.md index a9dddf5..0f6bcd5 100644 --- a/ARCHITECTURE.md +++ b/ARCHITECTURE.md @@ -29,14 +29,10 @@ RsChat uses a hybrid streaming architecture that provides both real-time perform ### Key Components -#### 1. LlmStreamWriter (`server/src/stream/llm_writer.rs`) +#### 1. LlmClientStreamer (`server/src/stream/streamer.rs`) The core component that processes LLM provider streams and manages Redis stream output. -**Key Features:** -- **Batching**: Accumulates chunks from the provider stream, up to a max length or timeout, and adds them to the Redis stream -- **Background Pings**: Sends regular keepalive pings - #### 2. Redis and SSE Stream Structure **Redis Key for Chat Streams**: `user:{user_id}:chat:{session_id}` @@ -68,8 +64,7 @@ Stream End → Database Save → Redis DEL #### Cross-Instance Support - Redis streams provide shared state across server instances -- Background ping tasks maintain stream liveness -- Stream cancellation detected via Redis XADD failures +- Stream cancellation detection ## Data Flow @@ -77,23 +72,21 @@ Stream End → Database Save → Redis DEL ``` Client → POST /api/chat/{session_id} → Send request to LLM Provider - → LLM response received, streamed to Redis with the `LlmStreamWriter` - → GET /api/chat/{session_id}/stream to connect to the stream and stream the response + → LLM response received, streamed to Redis with the `LlmClientStreamer` ``` ### 2. Stream Processing ``` LLM Chunk → Process text, tool calls, usage, and error chunks → Batching Logic - → Redis XADD (if conditions met) - → Client(s) receive the new chunks + → Add chunks to Redis via `tinistream` service + → Client(s) receive the new chunks from `tinistream` ``` ### 3. Stream Completion ``` LLM End → Final Database Save - → Redis Stream End Event - → Redis Stream Cleanup + → `tinistream` End Event → SSE Connection Close ``` diff --git a/README.md b/README.md index ecd340c..928e5b4 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,6 @@ # RsChat 🤖💬 -A fast, secure, self-hostable chat application built with Rust, TypeScript, and React. Chat with multiple AI providers using your own API keys, with real-time streaming built-in. +A lightweight, secure, open-source, self-hostable chat application built with Rust, TypeScript, and React. Stream chats with multiple AI providers using your own API keys. Demo link: https://rs-chat.fly.dev/ (⚠️ This is a demo - don't expect your account/chats to be there when you come back. It may intermittently delete all data. Please also don't enter any sensitive information or confidential data) @@ -11,21 +11,20 @@ Demo link: https://rs-chat.fly.dev/ (⚠️ This is a demo - don't expect your a - **Multiple AI Providers**: Chat with AI models from OpenAI, Anthropic, and OpenRouter - **Streaming**: Streams responses using SSE (Server-Sent Events) - **Concurrent Streaming**: Seamlessly switch between multiple AI conversations streamed at the same time -- **Resumable Conversations**: Resume the conversation if your connection is lost or the page is refreshed +- **Resilient Streams**: Streaming continues if your connection is lost or the page is refreshed - **Code Highlighting**: Beautiful syntax highlighting for code blocks using [`rehype-highlight`](https://github.com/rehypejs/rehype-highlight) - **Dark Mode**: Dark/light theme support - **Responsive Design**: Mobile-friendly layout - **Search Chats**: Full-text search of chat session titles and messages - **Fast and Memory Efficient**: Rust backend using the [Rocket framework](https://rocket.rs/) - **Users & Authentication**: Login via OAuth providers (Google, GitHub, etc.), custom OIDC, and SSO header authentication -- **API Key Access and OpenAPI Docs**: API key access and documentation at `/api/docs` for developers to integrate with RsChat +- **Documented API**: API key access and documentation at `/api/docs` for developers to integrate with RsChat - **Fully Type-Safe**: End-to-end type safety with auto-generated client from OpenAPI spec ### ⚡ Convenience Features - **Smart Titles**: Auto-generation of chat titles -- **Smart Scrolling**: Auto-scroll during streaming and when opening previous chats -- **Secure Key Storage**: Your API keys are saved and encrypted +- **Auto Scrolling**: Auto-scroll during streaming and when opening previous chats ## 🏗️ Architecture @@ -45,14 +44,17 @@ rs-chat/ │ │ ├── auth/ # Authentication services │ │ ├── db/ # Database models and services │ │ ├── provider/ # AI provider integrations -│ │ ├── utils/ # Utility functions -│ │ ├── config.rs # Reading configuration / env variables +│ │ ├── storage/ # File storage services +│ │ ├── stream/ # Streaming utilities +│ │ ├── tools/ # AI chat tools +│ │ ├── utils/ # Other utilities +│ │ ├── config.rs # Configuration / environment variables │ │ ├── lib.rs # Server setup │ │ ├── main.rs # Server entry point │ │ └── ... # Other modules │ ├── migrations/ # Database migrations │ └── Cargo.toml # Rust dependencies -├── web/ # Vite / React frontend +├── web/ # Vite / React / TanStack Router frontend │ ├── src/ │ │ ├── components/ # React components │ │ ├── routes/ # TanStack Router routes @@ -90,9 +92,9 @@ Your API keys are encrypted and stored in the database. cd rs-chat ``` -2. **Start development databases** +2. **Start development services** ```bash - docker compose up -d db redis + docker compose up -d db redis stream ``` 3. **Set up the backend** @@ -145,33 +147,47 @@ You'll need an environment with PostgreSQL and Redis (or Redis-compatible databa services: rschat: image: ghcr.io/fa-sharp/rs-chat:latest - # ports: - # - "8080:8080" + ports: + - "8080:8080" environment: RUST_LOG: warn # 'info' or 'debug' for more logs RS_CHAT_SERVER_ADDRESS: https://mydomain.com # where you're hosting the app RS_CHAT_DATABASE_URL: postgres://user:pass@mypostgres/mydb # Your PostgreSQL URL RS_CHAT_REDIS_URL: redis://myredis:6379 # Your Redis URL RS_CHAT_SECRET_KEY: your-secret-key-for-encryption # 64-character hex string + RS_CHAT_TINISTREAM_URL: http://tinistream:8081 + RS_CHAT_TINISTREAM_API_KEY: tinistream-api-key # API key for the tinistream service + ## For GitHub login: callback URL should be {your_server_address}/api/auth/login/github/callback # RS_CHAT_GITHUB_CLIENT_ID: your-github-client-id # RS_CHAT_GITHUB_CLIENT_SECRET: your-github-client-secret ## Similar config for other OAuth providers - see server/src/auth/oauth/ folder # RS_CHAT_DISCORD_CLIENT_ID: your-discord-client-id # ... + ## For SSO header auth - see server/src/auth/sso_header.rs for all config options # RS_CHAT_SSO_HEADER_ENABLED: true # RS_CHAT_SSO_USERNAME_HEADER: X-Remote-User # ... + ## For running code on a remote Docker host # DOCKER_HOST: tcp://remote-docker-host:port # DOCKER_TLS_VERIFY: 1 # DOCKER_CERT_PATH: /certs volumes: - ## For running code on local Docker host - # - /var/run/docker.sock:/var/run/docker.sock:ro - ## Certificates for remote Docker host - # - ./path/to/certs:/certs + # - /var/run/docker.sock:/var/run/docker.sock:ro # To run code on local Docker + # - ./path/to/certs:/certs # Certificates for remote Docker host + tinistream: + image: ghcr.io/fa-sharp/tinistream:latest + container_name: tinistream + ports: + - "8081:8081" + environment: + STREAMER_PORT: 8081 + STREAMER_SERVER_ADDRESS: http://localhost:8081 + STREAMER_REDIS_URL: redis://myredis:6379 # Your Redis URL + STREAMER_API_KEY: tinistream-api-key # should match the RS_CHAT_TINISTREAM_API_KEY above + STREAMER_SECRET_KEY: your-secret-key-for-encryption # 64-character hex string ``` ## 🔒 Security & Privacy diff --git a/server/.env.example b/server/.env.example index 69f7954..0b31db9 100644 --- a/server/.env.example +++ b/server/.env.example @@ -8,10 +8,15 @@ RS_CHAT_GITHUB_CLIENT_SECRET=your_github_client_secret_here # Generate a 64-character hex key for encryption # You can generate one with: openssl rand -hex 32 -RS_CHAT_SECRET_KEY=hex-secret-key-for-encryption-change-this +RS_CHAT_SECRET_KEY=64-character-hex-secret-key-change-this # Local data directory RS_CHAT_DATA_DIR=.local +# Tinistream service +RS_CHAT_TINISTREAM_API_KEY=api-key-123 +STREAMER_API_KEY=api-key-123 +STREAMER_SECRET_KEY=64-character-hex-secret-key-change-this + # Postgres URL for running migrations via the Diesel CLI DATABASE_URL=postgres://postgres:postgres@localhost/postgres From d5a6fab27f72edd8ac4fddf41c90c19b4b6e6cb1 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 17 Feb 2026 21:44:19 -0500 Subject: [PATCH 013/111] server: update tinistream client --- docker-compose.yml | 2 +- server/Cargo.lock | 215 +++++++++++++++++--------------- server/Cargo.toml | 8 +- server/src/stream/tinistream.rs | 2 +- 4 files changed, 119 insertions(+), 108 deletions(-) diff --git a/docker-compose.yml b/docker-compose.yml index 285da90..bcdb550 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -20,7 +20,7 @@ services: - redis_data:/data stream: - image: ghcr.io/fa-sharp/tinistream:0.1.9 + image: ghcr.io/fa-sharp/tinistream:0.1.10 platform: linux/amd64 container_name: tinistream ports: diff --git a/server/Cargo.lock b/server/Cargo.lock index f4f21b9..57931a2 100644 --- a/server/Cargo.lock +++ b/server/Cargo.lock @@ -164,9 +164,9 @@ dependencies = [ [[package]] name = "async-tungstenite" -version = "0.31.0" +version = "0.32.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ee88b4c88ac8c9ea446ad43498955750a4bbe64c4392f21ccfe5d952865e318f" +checksum = "8acc405d38be14342132609f06f02acaf825ddccfe76c4824a69281e0458ebd4" dependencies = [ "atomic-waker", "futures-core", @@ -391,12 +391,6 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" -[[package]] -name = "cfg_aliases" -version = "0.2.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" - [[package]] name = "chat-rs-api" version = "0.7.0" @@ -1065,6 +1059,21 @@ version = "0.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" +[[package]] +name = "foreign-types" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f6f339eb8adc052cd2ca78910fda869aefa38d22d5cb648e6485e4d3fc06f3b1" +dependencies = [ + "foreign-types-shared", +] + +[[package]] +name = "foreign-types-shared" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "00b0228411908ca8685dba7fc2cdd70ec9990a6e753e89b6ac91a84c40fbaf4b" + [[package]] name = "form_urlencoded" version = "1.2.2" @@ -1252,10 +1261,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0" dependencies = [ "cfg-if", - "js-sys", "libc", "wasi 0.11.1+wasi-snapshot-preview1", - "wasm-bindgen", ] [[package]] @@ -1265,11 +1272,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" dependencies = [ "cfg-if", - "js-sys", "libc", "r-efi", "wasip2", - "wasm-bindgen", ] [[package]] @@ -1546,13 +1551,28 @@ dependencies = [ "hyper 1.8.1", "hyper-util", "rustls 0.23.36", - "rustls-native-certs 0.8.3", "rustls-pki-types", "tokio", "tokio-rustls 0.26.4", "tower-service", ] +[[package]] +name = "hyper-tls" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "70206fc6890eaca9fde8a0bf71caa2ddfc9fe045ac9e5c70df101a7dbde866e0" +dependencies = [ + "bytes", + "http-body-util", + "hyper 1.8.1", + "hyper-util", + "native-tls", + "tokio", + "tokio-native-tls", + "tower-service", +] + [[package]] name = "hyper-util" version = "0.1.20" @@ -1907,12 +1927,6 @@ dependencies = [ "tracing-subscriber", ] -[[package]] -name = "lru-slab" -version = "0.1.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "112b39cec0b298b6c1999fee3e31427f74f676e4cb9879ed1a121b43661a4154" - [[package]] name = "matchers" version = "0.2.0" @@ -2001,6 +2015,23 @@ dependencies = [ "version_check", ] +[[package]] +name = "native-tls" +version = "0.2.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d5d26952a508f321b4d3d2e80e78fc2603eaefcdf0c30783867f19586518bdc" +dependencies = [ + "libc", + "log", + "openssl", + "openssl-probe 0.2.1", + "openssl-sys", + "schannel", + "security-framework 3.6.0", + "security-framework-sys", + "tempfile", +] + [[package]] name = "nom" version = "7.1.3" @@ -2157,6 +2188,32 @@ version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" +[[package]] +name = "openssl" +version = "0.10.75" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08838db121398ad17ab8531ce9de97b244589089e290a384c900cb9ff7434328" +dependencies = [ + "bitflags", + "cfg-if", + "foreign-types", + "libc", + "once_cell", + "openssl-macros", + "openssl-sys", +] + +[[package]] +name = "openssl-macros" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a948666b637a0f465e8564c73e89d4dde00d72d4d473cc972f390fc3dcee7d9c" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.116", +] + [[package]] name = "openssl-probe" version = "0.1.6" @@ -2169,6 +2226,18 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe" +[[package]] +name = "openssl-sys" +version = "0.9.111" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "82cab2d520aa75e3c58898289429321eb788c3106963d0dc886ec7a5f4adc321" +dependencies = [ + "cc", + "libc", + "pkg-config", + "vcpkg", +] + [[package]] name = "outref" version = "0.5.2" @@ -2400,9 +2469,9 @@ dependencies = [ [[package]] name = "progenitor-client" -version = "0.11.2" +version = "0.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "71a0beb939758f229cbae70a4889c7c76a4ac0e90f0b1e7ae9b4636a927d1018" +checksum = "ffab7b358944dba033a7b324e7558e66e6bcb1fb4705cf57f26fd5092bcae630" dependencies = [ "bytes", "futures-core", @@ -2413,61 +2482,6 @@ dependencies = [ "serde_urlencoded", ] -[[package]] -name = "quinn" -version = "0.11.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b9e20a958963c291dc322d98411f541009df2ced7b5a4f2bd52337638cfccf20" -dependencies = [ - "bytes", - "cfg_aliases", - "pin-project-lite", - "quinn-proto", - "quinn-udp", - "rustc-hash", - "rustls 0.23.36", - "socket2 0.6.2", - "thiserror", - "tokio", - "tracing", - "web-time", -] - -[[package]] -name = "quinn-proto" -version = "0.11.13" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f1906b49b0c3bc04b5fe5d86a77925ae6524a19b816ae38ce1e426255f1d8a31" -dependencies = [ - "bytes", - "getrandom 0.3.4", - "lru-slab", - "rand 0.9.2", - "ring", - "rustc-hash", - "rustls 0.23.36", - "rustls-pki-types", - "slab", - "thiserror", - "tinyvec", - "tracing", - "web-time", -] - -[[package]] -name = "quinn-udp" -version = "0.5.14" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "addec6a0dcad8a8d96a771f815f0eaf55f9d1805756410b39f5fa81332574cbd" -dependencies = [ - "cfg_aliases", - "libc", - "once_cell", - "socket2 0.6.2", - "tracing", - "windows-sys 0.60.2", -] - [[package]] name = "quote" version = "1.0.44" @@ -2639,9 +2653,9 @@ checksum = "a96887878f22d7bad8a3b6dc5b7440e0ada9a245242924394987b21cf2210a4c" [[package]] name = "reqwest" -version = "0.12.28" +version = "0.13.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147" +checksum = "ab3f43e3283ab1488b624b44b0e988d0acea0b3214e694730a055cb6b2efa801" dependencies = [ "base64 0.22.1", "bytes", @@ -2651,22 +2665,20 @@ dependencies = [ "http-body 1.0.1", "http-body-util", "hyper 1.8.1", - "hyper-rustls 0.27.7", + "hyper-tls", "hyper-util", "js-sys", "log", + "native-tls", "percent-encoding", "pin-project-lite", - "quinn", - "rustls 0.23.36", - "rustls-native-certs 0.8.3", "rustls-pki-types", "serde", "serde_json", "serde_urlencoded", "sync_wrapper", "tokio", - "tokio-rustls 0.26.4", + "tokio-native-tls", "tokio-util", "tower", "tower-http", @@ -2680,9 +2692,9 @@ dependencies = [ [[package]] name = "reqwest-websocket" -version = "0.5.1" +version = "0.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cd5f79b25f7f17a62cc9337108974431a66ae5a723ac0d9fe78ac1cce2027720" +checksum = "7705b649c3b66b85c4e9c304a6898b1ae3eecb880c474720ebf925e4a932ae02" dependencies = [ "async-tungstenite", "bytes", @@ -2963,7 +2975,6 @@ version = "1.14.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "be040f8b0a225e40375822a563fa9524378b9d63112f53e19ffff34df5d33fdd" dependencies = [ - "web-time", "zeroize", ] @@ -3519,8 +3530,8 @@ dependencies = [ [[package]] name = "tinistream-client" -version = "0.1.7" -source = "git+https://github.com/fa-sharp/tinistream?rev=c37e41d#c37e41dea494c82f06507387cedecb5c3737df47" +version = "0.1.10" +source = "git+https://github.com/fa-sharp/tinistream?rev=f25144c#f25144c1bdbee827d6033606b94a8aa1ae6eb5a7" dependencies = [ "bytes", "futures-core", @@ -3582,6 +3593,16 @@ dependencies = [ "syn 2.0.116", ] +[[package]] +name = "tokio-native-tls" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbae76ab933c85776efabc971569dd6119c580d8f5d448769dec1764bf796ef2" +dependencies = [ + "native-tls", + "tokio", +] + [[package]] name = "tokio-postgres" version = "0.7.16" @@ -3852,9 +3873,9 @@ checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" [[package]] name = "tungstenite" -version = "0.27.0" +version = "0.28.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "eadc29d668c91fcc564941132e17b28a7ceb2f3ebf0b9dae3e03fd7a6748eb0d" +checksum = "8628dcc84e5a09eb3d8423d6cb682965dea9133204e8fb3efee74c2a0c259442" dependencies = [ "bytes", "data-encoding", @@ -4158,9 +4179,9 @@ dependencies = [ [[package]] name = "wasm-streams" -version = "0.4.2" +version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "15053d8d85c7eccdbefef60f06769760a563c7f0a9d6902a13d35c7800b0ad65" +checksum = "9d1ec4f6517c9e11ae630e200b2b65d193279042e28edd4a2cda233e46670bbb" dependencies = [ "futures-util", "js-sys", @@ -4191,16 +4212,6 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "web-time" -version = "1.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5a6580f308b1fad9207618087a65c04e7a10bc77e02c8e84e9b00dd4b12fa0bb" -dependencies = [ - "js-sys", - "wasm-bindgen", -] - [[package]] name = "whoami" version = "2.1.1" diff --git a/server/Cargo.toml b/server/Cargo.toml index 658c70e..b35e121 100644 --- a/server/Cargo.toml +++ b/server/Cargo.toml @@ -36,12 +36,12 @@ fred = { version = "10.1.0", default-features = false, features = [ hex = "0.4.3" jsonschema = { version = "0.30.0", default-features = false } rand = "0.9.2" -reqwest = { version = "0.12.28", default-features = false, features = [ +reqwest = { version = "0.13.2", default-features = false, features = [ "json", - "rustls-tls-native-roots", + "native-tls-no-alpn", "stream", ] } -reqwest-websocket = { version = "0.5.1", features = ["json"] } +reqwest-websocket = { version = "0.6.0", features = ["json"] } rocket = { version = "0.5.1", features = ["json", "uuid"] } rocket_flex_session = { version = "0.2.0", features = [ "redis_fred", @@ -54,7 +54,7 @@ serde = { version = "1.0.228" } serde_json = "1.0.149" subst = { version = "0.3.8", features = ["json"] } thiserror = "2.0.18" -tinistream-client = { git = "https://github.com/fa-sharp/tinistream", rev = "c37e41d" } +tinistream-client = { git = "https://github.com/fa-sharp/tinistream", rev = "f25144c" } tokio = { version = "1.49.0" } tokio-stream = "0.1.18" tokio-util = { version = "0.7.18", features = ["io"] } diff --git a/server/src/stream/tinistream.rs b/server/src/stream/tinistream.rs index e362fa1..5f10b98 100644 --- a/server/src/stream/tinistream.rs +++ b/server/src/stream/tinistream.rs @@ -1,4 +1,4 @@ -use reqwest_websocket::{RequestBuilderExt, WebSocket}; +use reqwest_websocket::{Upgrade, WebSocket}; use tinistream_client::{types::*, Client, ClientEventsExt, ClientInfo, ClientStreamExt, Error}; /// A client for interacting with the tinistream API. From 646edb7a8cfeec0d6d9b53afca66f77c1d46d43f Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 17 Feb 2026 23:37:22 -0500 Subject: [PATCH 014/111] web: allow attaching files for all models --- web/src/components/chat/ChatMessageInput.tsx | 42 +++++++++----------- web/src/hooks/useChatInputState.tsx | 5 +-- 2 files changed, 20 insertions(+), 27 deletions(-) diff --git a/web/src/components/chat/ChatMessageInput.tsx b/web/src/components/chat/ChatMessageInput.tsx index 84bb4fe..973288f 100644 --- a/web/src/components/chat/ChatMessageInput.tsx +++ b/web/src/components/chat/ChatMessageInput.tsx @@ -208,29 +208,25 @@ export default memo(function ChatMessageInput({ onToggleExternalApiTool={onToggleExternalApiTool} /> )} - {(!selectedModel || selectedModel.attachment) && ( - <> - - {files.length > 0 && ( -
- - {files.map((file) => file.path).join(", ")} -
- )} - {uploadingFiles.length > 0 && ( -
- - Uploading... -
- )} - + + {files.length > 0 && ( +
+ + {files.map((file) => file.path).join(", ")} +
+ )} + {uploadingFiles.length > 0 && ( +
+ + Uploading... +
)} )} diff --git a/web/src/hooks/useChatInputState.tsx b/web/src/hooks/useChatInputState.tsx index f777dff..bc83a79 100644 --- a/web/src/hooks/useChatInputState.tsx +++ b/web/src/hooks/useChatInputState.tsx @@ -166,10 +166,7 @@ export const useChatInputState = ({ modalities: selectedModel?.modalities?.output, }, tools: selectedModel?.tool_call ? toolInput : undefined, - files: - selectedModel?.attachment && files.length > 0 - ? files.map((file) => file.id) - : undefined, + files: files.length > 0 ? files.map((file) => file.id) : undefined, }); formRef.current?.reset(); }, [ From 39b42d78b3dbd227812fefcce3175057dbc3c02e Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 17 Feb 2026 23:37:25 -0500 Subject: [PATCH 015/111] Update utils.rs --- server/src/provider/utils.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/server/src/provider/utils.rs b/server/src/provider/utils.rs index 713a66c..0f18a5b 100644 --- a/server/src/provider/utils.rs +++ b/server/src/provider/utils.rs @@ -12,7 +12,7 @@ use crate::provider::LlmStreamError; /// Create a data URI pub fn create_data_uri(content_type: &str, b64_string: &str) -> String { - format!("data:{};base64,{}", content_type, b64_string) + format!("data:{content_type};base64,{b64_string}") } /// Get a stream of deserialized events from a provider SSE stream. From 361c7f434f206a6238c296687e4651811284ae39 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 18 Feb 2026 03:26:05 -0500 Subject: [PATCH 016/111] web: sort models --- web/src/components/ProviderManager.tsx | 2 +- .../chat/settings/ChatModelSelect.tsx | 57 ++++++++++--------- web/src/lib/api/provider.ts | 1 + 3 files changed, 31 insertions(+), 29 deletions(-) diff --git a/web/src/components/ProviderManager.tsx b/web/src/components/ProviderManager.tsx index 8ec08ff..6ef52b2 100644 --- a/web/src/components/ProviderManager.tsx +++ b/web/src/components/ProviderManager.tsx @@ -78,7 +78,7 @@ const PROVIDERS: Record = { baseUrl: "https://openrouter.ai/api/v1", keyFormat: "sk-or-...", color: "bg-blue-100 dark:bg-blue-900 border-blue-300 dark:border-blue-700", - defaultModel: "moonshotai/kimi-k2.5", + defaultModel: "openai/gpt-4o-mini", }, ollama: { name: "Ollama", diff --git a/web/src/components/chat/settings/ChatModelSelect.tsx b/web/src/components/chat/settings/ChatModelSelect.tsx index c1c5cb5..1e7b02b 100644 --- a/web/src/components/chat/settings/ChatModelSelect.tsx +++ b/web/src/components/chat/settings/ChatModelSelect.tsx @@ -1,4 +1,4 @@ -import { ChevronsUpDown, Eye, FileText, ImageIcon, Wrench } from "lucide-react"; +import { ChevronsUpDown, Eye, FileText, Wrench } from "lucide-react"; import React from "react"; import PopoverDrawer from "@/components/PopoverDrawer"; @@ -54,34 +54,35 @@ export default function ChatModelSelect({ No models found. - {models?.map((model) => ( - { - onSelect(model.id); - setOpen(false); - }} - > -
(a.id === currentModelId ? -1 : 0)) + .map((model) => ( + { + onSelect(model.id); + setOpen(false); + }} > - {model.name} - - {model.id} - -
-
- {model.tool_call && } - {model.modalities?.input.includes("image") && } - {model.modalities?.input.includes("pdf") && } - {model.modalities?.output.includes("image") && } -
-
- ))} +
+ {model.name} + + {model.id} + +
+
+ {model.tool_call && } + {model.modalities?.input.includes("image") && } + {model.modalities?.input.includes("pdf") && } +
+ + ))}
diff --git a/web/src/lib/api/provider.ts b/web/src/lib/api/provider.ts index ddb6aeb..ecd4beb 100644 --- a/web/src/lib/api/provider.ts +++ b/web/src/lib/api/provider.ts @@ -31,6 +31,7 @@ export const useProviderModels = (providerId?: number | null) => if (response.error) { throw new Error(response.error.message); } + response.data.sort((a, b) => a.name.localeCompare(b.name)); return response.data; }, }); From 871b13385e8ee9886716b700c62c228ce7ca7819 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Fri, 27 Feb 2026 06:04:52 -0500 Subject: [PATCH 017/111] server: add tinirun client dep --- server/Cargo.lock | 169 +++++++++++++++++++++++++++++++++++++++++++++- server/Cargo.toml | 1 + 2 files changed, 169 insertions(+), 1 deletion(-) diff --git a/server/Cargo.lock b/server/Cargo.lock index 57931a2..f962721 100644 --- a/server/Cargo.lock +++ b/server/Cargo.lock @@ -375,6 +375,12 @@ dependencies = [ "either", ] +[[package]] +name = "cargo-husky" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b02b629252fe8ef6460461409564e2c21d0c8e77e0944f3d189ff06c4e932ad" + [[package]] name = "cc" version = "1.2.56" @@ -425,6 +431,7 @@ dependencies = [ "serde_json", "subst", "thiserror", + "tinirun-client", "tinistream-client", "tokio", "tokio-stream", @@ -588,6 +595,16 @@ dependencies = [ "darling_macro 0.13.4", ] +[[package]] +name = "darling" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc7f46116c46ff9ab3eb1597a45688b6715c6e628b5c133e288e709a29bcb4ee" +dependencies = [ + "darling_core 0.20.11", + "darling_macro 0.20.11", +] + [[package]] name = "darling" version = "0.21.3" @@ -622,6 +639,20 @@ dependencies = [ "syn 1.0.109", ] +[[package]] +name = "darling_core" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d00b9596d185e565c2207a0b01f8bd1a135483d02d9b7b0a54b11da8d53412e" +dependencies = [ + "fnv", + "ident_case", + "proc-macro2", + "quote", + "strsim 0.11.1", + "syn 2.0.116", +] + [[package]] name = "darling_core" version = "0.21.3" @@ -660,6 +691,17 @@ dependencies = [ "syn 1.0.109", ] +[[package]] +name = "darling_macro" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc34b93ccb385b40dc71c6fceac4b2ad23662c7eeb248cf10d529b7e055b6ead" +dependencies = [ + "darling_core 0.20.11", + "quote", + "syn 2.0.116", +] + [[package]] name = "darling_macro" version = "0.21.3" @@ -2445,6 +2487,28 @@ dependencies = [ "syn 2.0.116", ] +[[package]] +name = "proc-macro-error-attr2" +version = "2.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "96de42df36bb9bba5542fe9f1a054b8cc87e172759a1868aa05c1f3acc89dfc5" +dependencies = [ + "proc-macro2", + "quote", +] + +[[package]] +name = "proc-macro-error2" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "11ec05c52be0a07b08061f7dd003e7d7092e0472bc731b4af7bb1ef876109802" +dependencies = [ + "proc-macro-error-attr2", + "proc-macro2", + "quote", + "syn 2.0.116", +] + [[package]] name = "proc-macro2" version = "1.0.106" @@ -2690,6 +2754,22 @@ dependencies = [ "web-sys", ] +[[package]] +name = "reqwest-streams" +version = "0.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3d2484a49257a16e13f0c5f760c6e65a231eb32fbe6f899f6caa6f9bc78e2799" +dependencies = [ + "async-trait", + "bytes", + "cargo-husky", + "futures", + "reqwest", + "serde", + "serde_json", + "tokio-util", +] + [[package]] name = "reqwest-websocket" version = "0.6.0" @@ -3029,7 +3109,7 @@ dependencies = [ "chrono", "dyn-clone", "indexmap 1.9.3", - "schemars_derive", + "schemars_derive 0.8.22", "serde", "serde_json", "uuid", @@ -3053,8 +3133,10 @@ version = "1.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a2b42f36aa1cd011945615b92222f6bf73c599a102a300334cd7f8dbeec726cc" dependencies = [ + "chrono", "dyn-clone", "ref-cast", + "schemars_derive 1.2.1", "serde", "serde_json", ] @@ -3071,6 +3153,18 @@ dependencies = [ "syn 2.0.116", ] +[[package]] +name = "schemars_derive" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7d115b50f4aaeea07e79c1912f645c7513d81715d0420f8bc77a18c6260b307f" +dependencies = [ + "proc-macro2", + "quote", + "serde_derive_internals", + "syn 2.0.116", +] + [[package]] name = "scoped-futures" version = "0.1.4" @@ -3191,6 +3285,7 @@ version = "1.0.149" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86" dependencies = [ + "indexmap 2.13.0", "itoa", "memchr", "serde", @@ -3254,9 +3349,22 @@ dependencies = [ "schemars 1.2.1", "serde_core", "serde_json", + "serde_with_macros", "time", ] +[[package]] +name = "serde_with_macros" +version = "3.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52a8e3ca0ca629121f70ab50f95249e5a6f925cc0f6ffe8256c45b728875706c" +dependencies = [ + "darling 0.21.3", + "proc-macro2", + "quote", + "syn 2.0.116", +] + [[package]] name = "sha1" version = "0.10.6" @@ -3528,6 +3636,35 @@ dependencies = [ "time-core", ] +[[package]] +name = "tinirun-client" +version = "0.1.1" +source = "git+https://github.com/fa-sharp/tinirun?rev=1413a49#1413a495a949d227cfec08f69afbe001987ded94" +dependencies = [ + "futures", + "reqwest", + "reqwest-streams", + "serde", + "serde_json", + "thiserror", + "tinirun-models", + "validator", +] + +[[package]] +name = "tinirun-models" +version = "0.1.1" +source = "git+https://github.com/fa-sharp/tinirun?rev=1413a49#1413a495a949d227cfec08f69afbe001987ded94" +dependencies = [ + "regex", + "schemars 1.2.1", + "serde", + "serde_json", + "serde_with", + "thiserror", + "validator", +] + [[package]] name = "tinistream-client" version = "0.1.10" @@ -4021,6 +4158,36 @@ dependencies = [ "vsimd", ] +[[package]] +name = "validator" +version = "0.20.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "43fb22e1a008ece370ce08a3e9e4447a910e92621bb49b85d6e48a45397e7cfa" +dependencies = [ + "idna", + "once_cell", + "regex", + "serde", + "serde_derive", + "serde_json", + "url", + "validator_derive", +] + +[[package]] +name = "validator_derive" +version = "0.20.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7df16e474ef958526d1205f6dda359fdfab79d9aa6d54bafcb92dcd07673dca" +dependencies = [ + "darling 0.20.11", + "once_cell", + "proc-macro-error2", + "proc-macro2", + "quote", + "syn 2.0.116", +] + [[package]] name = "valuable" version = "0.1.1" diff --git a/server/Cargo.toml b/server/Cargo.toml index b35e121..cdc8595 100644 --- a/server/Cargo.toml +++ b/server/Cargo.toml @@ -54,6 +54,7 @@ serde = { version = "1.0.228" } serde_json = "1.0.149" subst = { version = "0.3.8", features = ["json"] } thiserror = "2.0.18" +tinirun-client = { git = "https://github.com/fa-sharp/tinirun", rev = "1413a49" } tinistream-client = { git = "https://github.com/fa-sharp/tinistream", rev = "f25144c" } tokio = { version = "1.49.0" } tokio-stream = "0.1.18" From e9239f9687b5bb92611e829da796bacf6802f9f0 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Fri, 27 Feb 2026 06:05:14 -0500 Subject: [PATCH 018/111] server: add openai subtype for possible future use --- server/src/db/models/provider.rs | 12 ++++++++++++ server/src/provider/providers/openai/response.rs | 6 ++++-- 2 files changed, 16 insertions(+), 2 deletions(-) diff --git a/server/src/db/models/provider.rs b/server/src/db/models/provider.rs index 4dda182..f7dd3e2 100644 --- a/server/src/db/models/provider.rs +++ b/server/src/db/models/provider.rs @@ -14,6 +14,8 @@ pub struct ChatRsProvider { pub name: String, #[schemars(with = "ChatRsProviderType")] pub provider_type: String, + // #[schemars(with = "OpenaiSubtype")] + // pub openai_subtype: Option, #[serde(skip)] pub user_id: Uuid, pub default_model: String, @@ -52,6 +54,16 @@ pub enum ChatRsProviderType { Lorem, } +/// The subtype for OpenAI-compatible providers +#[derive(JsonSchema, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum OpenaiSubtype { + Openai, + Google, + OpenRouter, + LlmGateway, +} + impl TryFrom<&str> for ChatRsProviderType { type Error = LlmError; diff --git a/server/src/provider/providers/openai/response.rs b/server/src/provider/providers/openai/response.rs index 667a169..bad7fd7 100644 --- a/server/src/provider/providers/openai/response.rs +++ b/server/src/provider/providers/openai/response.rs @@ -157,8 +157,10 @@ pub struct OpenRouterImageData { pub struct OpenAIUsage { prompt_tokens: Option, completion_tokens: Option, + /// OpenRouter cost cost: Option, - // total_tokens: Option, + /// LLM Gateway cost + cost_usd_total: Option, } impl From for LlmUsage { @@ -166,7 +168,7 @@ impl From for LlmUsage { LlmUsage { input_tokens: usage.prompt_tokens, output_tokens: usage.completion_tokens, - cost: usage.cost, + cost: usage.cost.or(usage.cost_usd_total), } } } From 81d3dd6449a9a8cf2d869060579d73f6aa18447d Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Fri, 27 Feb 2026 06:24:59 -0500 Subject: [PATCH 019/111] server: add code runner to docker compose --- .docker/.gitkeep | 0 .gitignore | 1 + docker-compose.yml | 19 ++++++++++++++++++- 3 files changed, 19 insertions(+), 1 deletion(-) delete mode 100644 .docker/.gitkeep diff --git a/.docker/.gitkeep b/.docker/.gitkeep deleted file mode 100644 index e69de29..0000000 diff --git a/.gitignore b/.gitignore index 2391a64..32dd7e0 100644 --- a/.gitignore +++ b/.gitignore @@ -35,6 +35,7 @@ Desktop.ini *.pfx secrets.json config/secrets.yml +.docker/ .secrets/ # Dependency Directories diff --git a/docker-compose.yml b/docker-compose.yml index bcdb550..0458c2a 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -34,6 +34,22 @@ services: depends_on: - redis + runner: + image: ghcr.io/fa-sharp/tinirun:0.1.1 + container_name: tinirun + ports: + - "8082:8082" + environment: + RUNNER_HOST: 0.0.0.0 + RUNNER_PORT: 8082 + RUNNER_LOG_LEVEL: info + RUNNER_REDIS_URL: redis://redis:6379 + env_file: server/.env + depends_on: + - redis + volumes: + - /var/run/docker.sock:/var/run/docker.sock + rschat: build: context: . @@ -46,14 +62,15 @@ services: RS_CHAT_REDIS_URL: redis://redis:6379 RS_CHAT_DATA_DIR: /data RS_CHAT_TINISTREAM_URL: http://tinistream:8081 + RS_CHAT_RUNNER_URL: http://tinirun:8082 env_file: server/.env volumes: - - ./.docker:/certs - rschat_data:/data depends_on: - db - redis - stream + - runner volumes: postgres_data: From 6e0027c29a2a271d6a1e03607cf1e1afd26b20e9 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sat, 28 Feb 2026 19:59:04 -0500 Subject: [PATCH 020/111] server: implement tinirun as code runner, remove unneeded deps --- docker-compose.yml | 42 +- server/.env.example | 6 +- server/Cargo.lock | 294 +----------- server/Cargo.toml | 5 +- server/Rocket.toml | 1 + server/src/api/tool.rs | 2 +- server/src/auth/sso_header.rs | 37 +- server/src/config.rs | 7 + server/src/tools/system.rs | 8 +- server/src/tools/system/code_runner.rs | 35 +- server/src/tools/system/code_runner/docker.rs | 435 ------------------ .../tools/system/code_runner/dockerfiles.rs | 162 ------- .../src/tools/system/code_runner/tinirun.rs | 129 ++++++ .../chat/ChatStreamingToolCalls.tsx | 11 +- .../chat/messages/ChatMessageToolLogs.tsx | 37 +- .../chat/messages/ChatMessageToolResult.tsx | 2 +- .../components/ui/chat/chat-message-list.tsx | 1 - 17 files changed, 246 insertions(+), 968 deletions(-) delete mode 100644 server/src/tools/system/code_runner/docker.rs delete mode 100644 server/src/tools/system/code_runner/dockerfiles.rs create mode 100644 server/src/tools/system/code_runner/tinirun.rs diff --git a/docker-compose.yml b/docker-compose.yml index 0458c2a..47e605d 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -50,27 +50,27 @@ services: volumes: - /var/run/docker.sock:/var/run/docker.sock - rschat: - build: - context: . - ports: - - "8080:8080" - environment: - RUST_LOG: info - RS_CHAT_SERVER_ADDRESS: http://localhost:8080 - RS_CHAT_DATABASE_URL: postgres://postgres:postgres@postgres/postgres - RS_CHAT_REDIS_URL: redis://redis:6379 - RS_CHAT_DATA_DIR: /data - RS_CHAT_TINISTREAM_URL: http://tinistream:8081 - RS_CHAT_RUNNER_URL: http://tinirun:8082 - env_file: server/.env - volumes: - - rschat_data:/data - depends_on: - - db - - redis - - stream - - runner + # rschat: + # build: + # context: . + # ports: + # - "8080:8080" + # environment: + # RUST_LOG: info + # RS_CHAT_SERVER_ADDRESS: http://localhost:8080 + # RS_CHAT_DATABASE_URL: postgres://postgres:postgres@postgres/postgres + # RS_CHAT_REDIS_URL: redis://redis:6379 + # RS_CHAT_DATA_DIR: /data + # RS_CHAT_TINISTREAM_URL: http://tinistream:8081 + # RS_CHAT_TINIRUN_URL: http://tinirun:8082 + # env_file: server/.env + # volumes: + # - rschat_data:/data + # depends_on: + # - db + # - redis + # - stream + # - runner volumes: postgres_data: diff --git a/server/.env.example b/server/.env.example index 0b31db9..3c4cff7 100644 --- a/server/.env.example +++ b/server/.env.example @@ -13,10 +13,14 @@ RS_CHAT_SECRET_KEY=64-character-hex-secret-key-change-this # Local data directory RS_CHAT_DATA_DIR=.local -# Tinistream service +# Tinistream service (client streaming utility) RS_CHAT_TINISTREAM_API_KEY=api-key-123 STREAMER_API_KEY=api-key-123 STREAMER_SECRET_KEY=64-character-hex-secret-key-change-this +# Tinirun service (code runner) +RS_CHAT_TINIRUN_API_KEY=api-key-123 +RUNNER_API_KEY=api-key-123 + # Postgres URL for running migrations via the Diesel CLI DATABASE_URL=postgres://postgres:postgres@localhost/postgres diff --git a/server/Cargo.lock b/server/Cargo.lock index f962721..3bb2ec5 100644 --- a/server/Cargo.lock +++ b/server/Cargo.lock @@ -84,22 +84,6 @@ dependencies = [ "rustversion", ] -[[package]] -name = "astral-tokio-tar" -version = "0.5.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ec179a06c1769b1e42e1e2cbe74c7dcdb3d6383c838454d063eaac5bbb7ebbe5" -dependencies = [ - "filetime", - "futures-core", - "libc", - "portable-atomic", - "rustc-hash", - "tokio", - "tokio-stream", - "xattr", -] - [[package]] name = "async-io" version = "2.6.0" @@ -253,57 +237,6 @@ dependencies = [ "generic-array", ] -[[package]] -name = "bollard" -version = "0.19.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "87a52479c9237eb04047ddb94788c41ca0d26eaff8b697ecfbb4c32f7fdc3b1b" -dependencies = [ - "base64 0.22.1", - "bollard-stubs", - "bytes", - "futures-core", - "futures-util", - "hex", - "home", - "http 1.4.0", - "http-body-util", - "hyper 1.8.1", - "hyper-named-pipe", - "hyper-rustls 0.27.7", - "hyper-util", - "hyperlocal", - "log", - "pin-project-lite", - "rustls 0.23.36", - "rustls-native-certs 0.8.3", - "rustls-pemfile 2.2.0", - "rustls-pki-types", - "serde", - "serde_derive", - "serde_json", - "serde_repr", - "serde_urlencoded", - "thiserror", - "tokio", - "tokio-util", - "tower-service", - "url", - "winapi", -] - -[[package]] -name = "bollard-stubs" -version = "1.49.1-rc.28.4.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5731fe885755e92beff1950774068e0cae67ea6ec7587381536fca84f1779623" -dependencies = [ - "serde", - "serde_json", - "serde_repr", - "serde_with", -] - [[package]] name = "bon" version = "3.9.0" @@ -402,10 +335,7 @@ name = "chat-rs-api" version = "0.7.0" dependencies = [ "aes-gcm", - "astral-tokio-tar", "base64 0.22.1", - "bollard", - "bon", "chrono", "const_format", "diesel", @@ -755,7 +685,6 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cc3dc5ad92c2e2d1c193bbbbdf2ea477cb81331de4f3103f267ca18368b988c4" dependencies = [ "powerfmt", - "serde_core", ] [[package]] @@ -1052,17 +981,6 @@ dependencies = [ "version_check", ] -[[package]] -name = "filetime" -version = "0.2.27" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f98844151eee8917efc50bd9e8318cb963ae8b297431495d3f758616ea5c57db" -dependencies = [ - "cfg-if", - "libc", - "libredox", -] - [[package]] name = "find-msvc-tools" version = "0.1.9" @@ -1430,15 +1348,6 @@ dependencies = [ "digest", ] -[[package]] -name = "home" -version = "0.5.12" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cc627f471c528ff0c4a49e1d5e60450c8f6461dd6d10ba9dcd3a61d3dff7728d" -dependencies = [ - "windows-sys 0.61.2", -] - [[package]] name = "http" version = "0.2.12" @@ -1543,7 +1452,6 @@ dependencies = [ "http 1.4.0", "http-body 1.0.1", "httparse", - "httpdate", "itoa", "pin-project-lite", "pin-utils", @@ -1552,21 +1460,6 @@ dependencies = [ "want", ] -[[package]] -name = "hyper-named-pipe" -version = "0.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "73b7d8abf35697b81a825e386fc151e0d503e8cb5fcb93cc8669c376dfd6f278" -dependencies = [ - "hex", - "hyper 1.8.1", - "hyper-util", - "pin-project-lite", - "tokio", - "tower-service", - "winapi", -] - [[package]] name = "hyper-rustls" version = "0.24.2" @@ -1577,26 +1470,10 @@ dependencies = [ "http 0.2.12", "hyper 0.14.32", "log", - "rustls 0.21.12", - "rustls-native-certs 0.6.3", - "tokio", - "tokio-rustls 0.24.1", -] - -[[package]] -name = "hyper-rustls" -version = "0.27.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e3c93eb611681b207e1fe55d5a71ecf91572ec8a6705cdb6857f7d8d5242cf58" -dependencies = [ - "http 1.4.0", - "hyper 1.8.1", - "hyper-util", - "rustls 0.23.36", - "rustls-pki-types", + "rustls", + "rustls-native-certs", "tokio", - "tokio-rustls 0.26.4", - "tower-service", + "tokio-rustls", ] [[package]] @@ -1638,21 +1515,6 @@ dependencies = [ "tracing", ] -[[package]] -name = "hyperlocal" -version = "0.9.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "986c5ce3b994526b3cd75578e62554abd09f0899d6206de48b3e96ab34ccc8c7" -dependencies = [ - "hex", - "http-body-util", - "hyper 1.8.1", - "hyper-util", - "pin-project-lite", - "tokio", - "tower-service", -] - [[package]] name = "iana-time-zone" version = "0.1.65" @@ -1924,7 +1786,6 @@ checksum = "3d0b95e02c851351f877147b7deea7b1afb1df71b63aa5f8270716e0c5720616" dependencies = [ "bitflags", "libc", - "redox_syscall 0.7.1", ] [[package]] @@ -2310,7 +2171,7 @@ checksum = "2621685985a2ebf1c516881c026032ac7deafcda1a2c9b7850dc81e3dfcb64c1" dependencies = [ "cfg-if", "libc", - "redox_syscall 0.5.18", + "redox_syscall", "smallvec", "windows-link", ] @@ -2407,12 +2268,6 @@ dependencies = [ "universal-hash", ] -[[package]] -name = "portable-atomic" -version = "1.13.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c33a9471896f1c69cecef8d20cbe2f7accd12527ce60845ff44c153bb2a21b49" - [[package]] name = "postgres-protocol" version = "0.6.10" @@ -2643,15 +2498,6 @@ dependencies = [ "bitflags", ] -[[package]] -name = "redox_syscall" -version = "0.7.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "35985aa610addc02e24fc232012c86fd11f14111180f902b67e2d5331f8ebf2b" -dependencies = [ - "bitflags", -] - [[package]] name = "ref-cast" version = "1.0.25" @@ -2925,7 +2771,7 @@ dependencies = [ "async-trait", "base64 0.21.7", "hyper 0.14.32", - "hyper-rustls 0.24.2", + "hyper-rustls", "log", "rand 0.8.5", "rocket", @@ -2962,12 +2808,6 @@ dependencies = [ "syn 1.0.109", ] -[[package]] -name = "rustc-hash" -version = "2.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "357703d41365b4b27c590e3ed91eabb1b663f07c4c084095e60cbed4362dff0d" - [[package]] name = "rustix" version = "1.1.3" @@ -2989,24 +2829,10 @@ checksum = "3f56a14d1f48b391359b22f731fd4bd7e43c97f3c50eee276f3aa09c94784d3e" dependencies = [ "log", "ring", - "rustls-webpki 0.101.7", + "rustls-webpki", "sct", ] -[[package]] -name = "rustls" -version = "0.23.36" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c665f33d38cea657d9614f766881e4d510e0eda4239891eea56b4cadcf01801b" -dependencies = [ - "once_cell", - "ring", - "rustls-pki-types", - "rustls-webpki 0.103.9", - "subtle", - "zeroize", -] - [[package]] name = "rustls-native-certs" version = "0.6.3" @@ -3014,23 +2840,11 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a9aace74cb666635c918e9c12bc0d348266037aa8eb599b5cba565709a8dff00" dependencies = [ "openssl-probe 0.1.6", - "rustls-pemfile 1.0.4", + "rustls-pemfile", "schannel", "security-framework 2.11.1", ] -[[package]] -name = "rustls-native-certs" -version = "0.8.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "612460d5f7bea540c490b2b6395d8e34a953e52b491accd6c86c8164c5932a63" -dependencies = [ - "openssl-probe 0.2.1", - "rustls-pki-types", - "schannel", - "security-framework 3.6.0", -] - [[package]] name = "rustls-pemfile" version = "1.0.4" @@ -3040,15 +2854,6 @@ dependencies = [ "base64 0.21.7", ] -[[package]] -name = "rustls-pemfile" -version = "2.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dce314e5fee3f39953d46bb63bb8a46d40c2f8fb7cc5a3b6cab2bde9721d6e50" -dependencies = [ - "rustls-pki-types", -] - [[package]] name = "rustls-pki-types" version = "1.14.0" @@ -3068,17 +2873,6 @@ dependencies = [ "untrusted", ] -[[package]] -name = "rustls-webpki" -version = "0.103.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d7df23109aa6c1567d1c575b9952556388da57401e4ace1d15f79eedad0d8f53" -dependencies = [ - "ring", - "rustls-pki-types", - "untrusted", -] - [[package]] name = "rustversion" version = "1.0.22" @@ -3115,18 +2909,6 @@ dependencies = [ "uuid", ] -[[package]] -name = "schemars" -version = "0.9.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4cd191f9397d57d581cddd31014772520aa448f65ef991055d7f61582c65165f" -dependencies = [ - "dyn-clone", - "ref-cast", - "serde", - "serde_json", -] - [[package]] name = "schemars" version = "1.2.1" @@ -3293,17 +3075,6 @@ dependencies = [ "zmij", ] -[[package]] -name = "serde_repr" -version = "0.1.20" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "175ee3e80ae9982737ca543e96133087cbd9a485eecc3bc4de9c1a37b47ea59c" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.116", -] - [[package]] name = "serde_spanned" version = "0.6.9" @@ -3343,9 +3114,6 @@ dependencies = [ "base64 0.22.1", "chrono", "hex", - "indexmap 1.9.3", - "indexmap 2.13.0", - "schemars 0.9.0", "schemars 1.2.1", "serde_core", "serde_json", @@ -3639,7 +3407,7 @@ dependencies = [ [[package]] name = "tinirun-client" version = "0.1.1" -source = "git+https://github.com/fa-sharp/tinirun?rev=1413a49#1413a495a949d227cfec08f69afbe001987ded94" +source = "git+https://github.com/fa-sharp/tinirun?rev=a60644d#a60644d36cdfc338a3268402a4bcabc33085245d" dependencies = [ "futures", "reqwest", @@ -3654,7 +3422,7 @@ dependencies = [ [[package]] name = "tinirun-models" version = "0.1.1" -source = "git+https://github.com/fa-sharp/tinirun?rev=1413a49#1413a495a949d227cfec08f69afbe001987ded94" +source = "git+https://github.com/fa-sharp/tinirun?rev=a60644d#a60644d36cdfc338a3268402a4bcabc33085245d" dependencies = [ "regex", "schemars 1.2.1", @@ -3772,17 +3540,7 @@ version = "0.24.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c28327cf380ac148141087fbfb9de9d7bd4e84ab5d2c28fbc911d753de8a7081" dependencies = [ - "rustls 0.21.12", - "tokio", -] - -[[package]] -name = "tokio-rustls" -version = "0.26.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61" -dependencies = [ - "rustls 0.23.36", + "rustls", "tokio", ] @@ -4392,28 +4150,6 @@ dependencies = [ "web-sys", ] -[[package]] -name = "winapi" -version = "0.3.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5c839a674fcd7a98952e593242ea400abe93992746761e38641405d28b00f419" -dependencies = [ - "winapi-i686-pc-windows-gnu", - "winapi-x86_64-pc-windows-gnu", -] - -[[package]] -name = "winapi-i686-pc-windows-gnu" -version = "0.4.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ac3b87c63620426dd9b991e5ce0329eff545bccbbb34f3be09ff6fb6ab51b7b6" - -[[package]] -name = "winapi-x86_64-pc-windows-gnu" -version = "0.4.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f" - [[package]] name = "windows" version = "0.48.0" @@ -4798,16 +4534,6 @@ version = "0.6.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9edde0db4769d2dc68579893f2306b26c6ecfbe0ef499b013d731b7b9247e0b9" -[[package]] -name = "xattr" -version = "1.6.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156" -dependencies = [ - "libc", - "rustix", -] - [[package]] name = "yansi" version = "1.0.1" diff --git a/server/Cargo.toml b/server/Cargo.toml index cdc8595..146984d 100644 --- a/server/Cargo.toml +++ b/server/Cargo.toml @@ -6,10 +6,7 @@ publish = false [dependencies] aes-gcm = "0.10.3" -astral-tokio-tar = "0.5.6" base64 = "0.22.1" -bollard = { version = "0.19.4", features = ["ssl"] } -bon = "3.9" chrono = { version = "0.4.43", features = ["serde"] } const_format = "0.2.35" diesel = { version = "2.3.6", features = [ @@ -54,7 +51,7 @@ serde = { version = "1.0.228" } serde_json = "1.0.149" subst = { version = "0.3.8", features = ["json"] } thiserror = "2.0.18" -tinirun-client = { git = "https://github.com/fa-sharp/tinirun", rev = "1413a49" } +tinirun-client = { git = "https://github.com/fa-sharp/tinirun", rev = "a60644d" } tinistream-client = { git = "https://github.com/fa-sharp/tinistream", rev = "f25144c" } tokio = { version = "1.49.0" } tokio-stream = "0.1.18" diff --git a/server/Rocket.toml b/server/Rocket.toml index 0cf5d9b..874f48f 100644 --- a/server/Rocket.toml +++ b/server/Rocket.toml @@ -8,3 +8,4 @@ server_address = "http://localhost:8000" database_url = "postgres://postgres:postgres@localhost/postgres" redis_url = "redis://localhost:6379" tinistream_url = "http://localhost:8081" +tinirun_url = "http://localhost:8082/api" diff --git a/server/src/api/tool.rs b/server/src/api/tool.rs index e2053b3..4265800 100644 --- a/server/src/api/tool.rs +++ b/server/src/api/tool.rs @@ -214,7 +214,7 @@ async fn execute_tool( let tool_result = match (system_tool, external_api_tool) { (Some(system_tool), None) => { system_tool - .build_executor(&mut db, &app_config, &message.session_id) + .build_executor(&mut db, &app_config, &http_client, &message.session_id) .validate_and_execute( &tool_call.tool_name, &tool_call.parameters, diff --git a/server/src/auth/sso_header.rs b/server/src/auth/sso_header.rs index 5df4b1f..ebe31dc 100644 --- a/server/src/auth/sso_header.rs +++ b/server/src/auth/sso_header.rs @@ -17,25 +17,34 @@ struct SsoHeaderConfig { /// Whether SSO header authentication is enabled sso_header_enabled: bool, /// Header for unique, identifying username (default: `Remote-User`) - sso_username_header: Option, + #[serde(default = "default_username_header")] + sso_username_header: String, /// Header for display name (default: `Remote-Name`) - sso_name_header: Option, + #[serde(default = "default_name_header")] + sso_name_header: String, /// Header for groups the user belongs to (default: `Remote-Groups`) - sso_groups_header: Option, + #[serde(default = "default_groups_header")] + sso_groups_header: String, /// If set, only users in this group will be allowed to access the app sso_user_group: Option, /// URL to redirect to in order to log out of the remote service sso_logout_url: Option, } +fn default_username_header() -> String { + "Remote-User".to_string() +} +fn default_name_header() -> String { + "Remote-Name".to_string() +} +fn default_groups_header() -> String { + "Remote-Groups".to_string() +} /// SSO header config added to Rocket state when enabled -#[derive(bon::Builder, Debug, Deserialize)] +#[derive(Debug, Deserialize)] pub struct SsoHeaderMergedConfig { - #[builder(into, default = "Remote-User")] pub username_header: String, - #[builder(into, default = "Remote-Name")] pub name_header: String, - #[builder(into, default = "Remote-Groups")] pub groups_header: String, pub user_group: Option, pub logout_url: Option, @@ -54,13 +63,13 @@ pub fn setup_sso_header_auth() -> AdHoc { match get_config_provider().extract::() { Ok(config) => { if config.sso_header_enabled { - let merged_config = SsoHeaderMergedConfig::builder() - .maybe_username_header(config.sso_username_header) - .maybe_name_header(config.sso_name_header) - .maybe_groups_header(config.sso_groups_header) - .maybe_user_group(config.sso_user_group) - .maybe_logout_url(config.sso_logout_url) - .build(); + let merged_config = SsoHeaderMergedConfig { + username_header: config.sso_username_header, + name_header: config.sso_name_header, + groups_header: config.sso_groups_header, + user_group: config.sso_user_group, + logout_url: config.sso_logout_url, + }; rocket::info!("SSO header auth: enabled! Config: {:?}", merged_config); rocket.manage(merged_config) } else { diff --git a/server/src/config.rs b/server/src/config.rs index 48df888..9a62bcf 100644 --- a/server/src/config.rs +++ b/server/src/config.rs @@ -18,16 +18,23 @@ pub struct AppConfig { pub static_path: Option, /// Local data directory (default: "/data") pub data_dir: Option, + /// Postgres Database URL pub database_url: String, /// Redis connection URL pub redis_url: String, /// Redis pool size (default: 4) pub redis_pool: Option, + /// Base URL of the tinistream API pub tinistream_url: String, /// API key for the tinistream API pub tinistream_api_key: String, + + /// Base URL of the tinirun API + pub tinirun_url: String, + /// API key for the tinirun API + pub tinirun_api_key: String, } /// Get the server configuration variables from Rocket diff --git a/server/src/tools/system.rs b/server/src/tools/system.rs index b197666..f9aacd0 100644 --- a/server/src/tools/system.rs +++ b/server/src/tools/system.rs @@ -137,11 +137,17 @@ impl<'a> ChatRsSystemTool { &'a self, db: &'a mut DbConnection, app_config: &'a AppConfig, + http_client: &'a reqwest::Client, session_id: &'a Uuid, ) -> Box { match &self.data { ChatRsSystemToolConfig::CodeRunner(config) => { - Box::new(code_runner::CodeRunner::new(config)) + let client = tinirun_client::TinirunClient::with_client( + http_client.to_owned(), + app_config.tinirun_url.clone(), + app_config.tinirun_api_key.clone(), + ); + Box::new(code_runner::CodeRunner::new(client, config)) } ChatRsSystemToolConfig::SystemInfo => { Box::new(system_info::SystemInfo::new(app_config)) diff --git a/server/src/tools/system/code_runner.rs b/server/src/tools/system/code_runner.rs index 5600587..0abf2b7 100644 --- a/server/src/tools/system/code_runner.rs +++ b/server/src/tools/system/code_runner.rs @@ -1,7 +1,3 @@ -mod docker; -mod dockerfiles; -use docker::{DockerExecutor, DockerExecutorOptions}; - use std::sync::LazyLock; use rocket::async_trait; @@ -20,14 +16,20 @@ use crate::{ utils::SenderWithLogging, }; +mod tinirun; +use tinirun::{TinirunExecutor, TinirunExecutorOptions}; + const CODE_RUNNER_NAME: &str = "code_runner"; const CODE_RUNNER_DESCRIPTION: &str = "Run code snippet in a sandboxed environment. \ - Temporary files can be written to the `$HOME` directory (must be created first). \ + Temporary files can be written to the /tmp directory and subdirectories. \ Other than that, it is a read-only environment."; const DEFAULT_TIMEOUT_SECONDS: u32 = 30; const DEFAULT_MEMORY_LIMIT_MB: u32 = 512; const DEFAULT_CPU_LIMIT: f32 = 0.5; +static CODE_RUNNER_INPUT_SCHEMA: LazyLock = + LazyLock::new(|| get_json_schema::()); + #[derive(Debug, Serialize, Deserialize, JsonSchema)] #[serde(deny_unknown_fields)] struct CodeRunnerInput { @@ -42,25 +44,22 @@ struct CodeRunnerInput { dependencies: Vec, // /// Whether to enable network access. Set to `true` only if the program needs to access the internet at runtime. // /// Network access is not needed for downloading dependencies. - // /// TODO: needs more safety precautions + // /// TODO: disabled for now because tinirun doesn't support this (security risk) // network: bool, } -static CODE_RUNNER_INPUT_SCHEMA: LazyLock = - LazyLock::new(|| get_json_schema::()); - -/// Tool to run code snippets in a sandboxed environment. -#[derive(Debug)] +/// Tool to run code snippets in a sandboxed environment. Uses the `tinirun` service. pub struct CodeRunner<'a> { + client: tinirun_client::TinirunClient, config: &'a CodeRunnerConfig, } impl<'a> CodeRunner<'a> { - pub fn new(config: &'a CodeRunnerConfig) -> Self { - CodeRunner { config } + pub fn new(client: tinirun_client::TinirunClient, config: &'a CodeRunnerConfig) -> Self { + CodeRunner { client, config } } } -#[derive(Debug, Serialize, Deserialize, JsonSchema)] +#[derive(Debug, Clone, Copy, Serialize, Deserialize, JsonSchema)] #[serde(rename_all = "lowercase")] enum CodeLanguage { Python, @@ -130,19 +129,19 @@ impl SystemTool for CodeRunner<'_> { sender: &SenderWithLogging, ) -> ToolResult<(String, ToolResponseFormat)> { let input = serde_json::from_value::(params)?; - let executor = DockerExecutor::new( + let executor = TinirunExecutor::new( + &self.client, input.language, - DockerExecutorOptions { + TinirunExecutorOptions { timeout_seconds: self.config.timeout_seconds, memory_limit_mb: self.config.memory_limit_mb, cpu_limit: self.config.cpu_limit, - network: false, }, ); - let tool_response = executor .execute(&input.code, &input.dependencies, sender) .await?; + Ok((tool_response, ToolResponseFormat::Markdown)) } } diff --git a/server/src/tools/system/code_runner/docker.rs b/server/src/tools/system/code_runner/docker.rs deleted file mode 100644 index 05dab29..0000000 --- a/server/src/tools/system/code_runner/docker.rs +++ /dev/null @@ -1,435 +0,0 @@ -use std::{sync::LazyLock, time::Duration}; - -use bollard::{ - body_try_stream, - container::{AttachContainerResults, LogOutput}, - models::{ContainerCreateBody, HostConfig, ResourcesUlimits}, - query_parameters::*, - Docker, -}; -use rocket::futures::StreamExt; -use tokio_util::io::ReaderStream; -use uuid::Uuid; - -use crate::{ - tools::{ - core::{ToolLog, ToolResult}, - system::code_runner::{ - dockerfiles::{get_dockerfile, get_dockerfile_info}, - CodeLanguage, - }, - ToolError, - }, - utils::SenderWithLogging, -}; - -static DOCKER: LazyLock> = - LazyLock::new(|| Docker::connect_with_defaults()); - -const GRACE_PERIOD_SECONDS: u32 = 5; - -pub struct DockerExecutor { - lang: CodeLanguage, - timeout_seconds: u32, - memory_limit_mb: u32, - cpu_limit: f32, - network: bool, - image_tag: String, - container_name: String, -} - -#[derive(Debug, Default)] -pub struct DockerExecutorOptions { - pub timeout_seconds: u32, - pub memory_limit_mb: u32, - pub cpu_limit: f32, - pub network: bool, -} - -impl DockerExecutor { - pub fn new(lang: CodeLanguage, options: DockerExecutorOptions) -> Self { - DockerExecutor { - lang, - timeout_seconds: options.timeout_seconds, - memory_limit_mb: options.memory_limit_mb, - cpu_limit: options.cpu_limit, - network: options.network, - image_tag: format!("code-runner-{}", Uuid::new_v4()), - container_name: format!("code-runner-{}", Uuid::new_v4()), - } - } - - pub async fn execute( - &self, - code: &str, - dependencies: &[String], - tx: &SenderWithLogging, - ) -> ToolResult { - let docker = DOCKER - .as_ref() - .inspect_err(|e| rocket::error!("Failed to initialize Docker client: {}", e)) - .map_err(|_| ToolError::ToolExecutionError("Failed to initialize Docker".into()))?; - docker - .ping() - .await - .inspect_err(|e| rocket::warn!("Failed to ping Docker daemon: {}", e)) - .map_err(|_| ToolError::ToolExecutionError("Couldn't connect to Docker".into()))?; - - // Run the code in a Docker container, returning early if the client disconnects - let result = tokio::select! { - result = self.run(docker, code, dependencies, &tx) => result, - _ = tx.closed() => Err(ToolError::Cancelled("client disconnected".to_string())) - }; - - // Cleanup container and image - if !tx.is_closed() { - send_log(tx, "Cleaning up...".into()).await; - docker_cleanup(docker, &self.container_name, &self.image_tag).await; - } else { - let container_name = self.container_name.clone(); - let image_tag = self.image_tag.clone(); - tokio::spawn(async move { - docker_cleanup(docker, &container_name, &image_tag).await; - }); - } - - result - } - - async fn run( - &self, - docker: &Docker, - code: &str, - dependencies: &[String], - tx: &SenderWithLogging, - ) -> ToolResult { - let (base_image, file_name, cmd) = get_dockerfile_info(&self.lang); - - // Check if base image exists locally, pull if needed - send_log(tx, format!("Checking base image '{base_image}'...")).await; - if docker.inspect_image(base_image).await.is_err() { - send_log(tx, format!("Pulling base image '{base_image}'...")).await; - let image_options = CreateImageOptionsBuilder::new() - .from_image(base_image) - .build(); - let mut pull_image_stream = docker.create_image(Some(image_options), None, None); - while let Some(result) = pull_image_stream.next().await { - match result { - Ok(mut response) => { - let status = response.status.unwrap_or_default(); - let progress_detail = response.progress_detail.take().unwrap_or_default(); - if let Some(progress) = response.progress { - send_debug(tx, format!("Pulling image: {status} {progress}")).await; - } else if let Some((current, total)) = - progress_detail.current.zip(progress_detail.total) - { - send_debug(tx, format!("Pulling image: {status} {current}/{total}")) - .await; - } - if let Some(error_detail) = response.error_detail { - send_error(tx, format!("Error pulling image: {:?}", error_detail)) - .await; - } - } - Err(err) => { - let message = format!("Error pulling image: {err}"); - send_error(tx, message.clone()).await; - return Err(ToolError::ToolExecutionError(message)); - } - } - } - } - - // Create tar archive with build context (Dockerfile and code files) - let (tar_writer, tar_reader) = tokio::io::duplex(8192); // 8KB buffer - let dockerfile = get_dockerfile(&self.lang); - let code = code.to_owned(); - send_log(tx, "Creating build context with 2 files...".into()).await; - - let tar_creation_task = tokio::spawn(async move { - let mut tar = tokio_tar::Builder::new(tar_writer); - for (path, content) in [("Dockerfile", dockerfile), (file_name, &code)] { - let mut header = tokio_tar::Header::new_gnu(); - header.set_size(content.len() as u64); - header.set_mode(0o644); - tar.append_data(&mut header, path, content.as_bytes()) - .await?; - } - tar.finish().await - }); - - // Build Docker image (streaming the build context tar file) - send_log(tx, format!("Building image '{}'...", self.image_tag)).await; - let build_options = BuildImageOptionsBuilder::new() - .buildargs(&[("DEPENDENCIES", self.build_dependency_string(dependencies))].into()) - .t(&self.image_tag) - .build(); - let mut build_stream = docker.build_image( - build_options, - None, - Some(body_try_stream(ReaderStream::new(tar_reader))), - ); - - let mut build_logs = String::new(); - let mut image_id = None; - while let Some(build_info_result) = build_stream.next().await { - match build_info_result { - Ok(info) => { - if let Some(id) = info.aux.and_then(|aux| aux.id) { - image_id = Some(id); - } - if let Some(stream) = info.stream { - build_logs.push_str(&format!("{stream}\n")); - send_debug(tx, stream).await; - } - if let Some(err) = info.error_detail.and_then(|e| e.message) { - build_logs.push_str(&format!("{err}\n")); - send_error(tx, format!("Error during build: {err}")).await; - } - } - Err(err) => { - build_logs.push_str(&format!("{err}\n")); - send_error(tx, format!("Error during build: {err}")).await; - } - } - } - if let Ok(Err(err)) = tar_creation_task.await { - let message = format!("Error while creating build context: {err}"); - send_error(tx, message).await; - } - if let Some(image_id) = image_id { - let message = format!("Built image '{}' with ID {}", self.image_tag, image_id); - send_log(tx, message).await; - } else { - let message = format!("Failed to build image '{}'", self.image_tag); - send_error(tx, message).await; - return Err(ToolError::ToolExecutionError(format!( - "Failed to build image '{}'. Build logs:\n\n{build_logs}", - self.image_tag - ))); - } - - // Create container with run command - let timeout_str = format!("{}s", self.timeout_seconds + GRACE_PERIOD_SECONDS); - let run_command = ["timeout", &timeout_str, "sh", "-c", &cmd]; - let container_body = ContainerCreateBody { - image: Some(self.image_tag.clone()), - cmd: Some(run_command.iter().map(|s| s.to_string()).collect()), - env: Some(vec!["HOME=/tmp/home".into()]), - network_disabled: Some(!self.network), - host_config: Some(HostConfig { - readonly_rootfs: Some(true), - tmpfs: Some([("/tmp".into(), "rw,noexec,nosuid,size=100m".into())].into()), - memory: Some((self.memory_limit_mb * 1024 * 1024).into()), - nano_cpus: Some((self.cpu_limit * 1000.0).round() as i64 * 1_000_000), - pids_limit: Some(50), - ulimits: Some(vec![ResourcesUlimits { - name: Some("nproc".into()), - soft: Some(50), - hard: Some(50), - }]), - cap_drop: Some(vec!["ALL".into()]), - security_opt: Some(vec!["no-new-privileges".into()]), - ..Default::default() - }), - ..Default::default() - }; - let container_options = CreateContainerOptionsBuilder::new() - .name(&self.container_name) - .build(); - match docker - .create_container(Some(container_options), container_body) - .await - { - Ok(res) => { - if !res.warnings.is_empty() { - let message = format!( - "⚠️ Warning while creating container '{}': {}", - self.container_name, - res.warnings.join(", ") - ); - send_log(tx, message).await; - } - } - Err(err) => { - let message = format!( - "Failed to create container '{}': {err}", - self.container_name - ); - send_error(tx, message.clone()).await; - return Err(ToolError::ToolExecutionError(message)); - } - }; - - // Spawn task to attach to container and capture logs/output - let attach_options = AttachContainerOptionsBuilder::new() - .stream(true) - .stdout(true) - .stderr(true) - .logs(true) - .build(); - let attached_container = match docker - .attach_container(&self.container_name, Some(attach_options)) - .await - { - Ok(container) => container, - Err(e) => { - let message = format!( - "Failed to attach to container '{}': {e}", - self.container_name - ); - send_error(tx, message.clone()).await; - return Err(ToolError::ToolExecutionError(message)); - } - }; - let output_tx = tx.clone(); - let output_timeout_secs = self.timeout_seconds + GRACE_PERIOD_SECONDS; - let container_output_task = tokio::spawn(async move { - let mut stdout = String::new(); - let mut stderr = String::new(); - let _ = tokio::time::timeout( - Duration::from_secs(output_timeout_secs.into()), - capture_container_output(attached_container, &mut stdout, &mut stderr, &output_tx), - ) - .await; - (stdout, stderr) - }); - - // Start container - send_log(tx, format!("Running command {run_command:?}...")).await; - if let Err(e) = docker - .start_container(&self.container_name, None::) - .await - { - let message = format!("Failed to start container '{}': {e}", self.container_name); - send_error(tx, message.clone()).await; - return Err(ToolError::ToolExecutionError(message)); - } - - // Wait for container to exit and get exit status - let container_exit_result = tokio::time::timeout( - Duration::from_secs(self.timeout_seconds.into()), - docker - .wait_container(&self.container_name, None::) - .next(), - ) - .await; - - // Process output and exit status - let (stdout, stderr) = container_output_task.await.unwrap_or_default(); - let output_text = format!("Output (stdout):\n\n{stdout}\n\nLogs (stderr):\n\n{stderr}\n"); - let output_markdown = - format!("## Output (stdout):\n```text\n{stdout}\n```\n## Logs (stderr):\n```text\n{stderr}\n```\n"); - match container_exit_result { - Err(_) => { - send_error(tx, "Code execution timed out".into()).await; - Err(ToolError::ToolExecutionError(format!( - "❌ Code execution timed out.\n\n{output_text}" - ))) - } - Ok(Some(wait_result)) => match wait_result { - Ok(_) => Ok(format!( - "✅ Code executed successfully!\n\n{output_markdown}" - )), - Err(err) => { - if let bollard::errors::Error::DockerContainerWaitError { code, .. } = err { - let message = format!("Code execution failed with exit status {code}"); - send_error(tx, message.clone()).await; - Err(ToolError::ToolExecutionError(format!( - "❌ {message}.\n\n{output_text}" - ))) - } else { - send_error(tx, "Code execution failed".into()).await; - Err(ToolError::ToolExecutionError(format!( - "❌ Code execution failed.\n\n{output_text}" - ))) - } - } - }, - Ok(None) => Ok(format!( - "Code executed with unknown exit status.\n\n{output_markdown}" - )), - } - } - - fn build_dependency_string(&self, dependencies: &[String]) -> String { - dependencies - .iter() - .filter_map(|d| { - let sanitized = self.sanitize_package_name(d); - if sanitized.trim().is_empty() { - None - } else { - Some(sanitized) - } - }) - .collect::>() - .join(" ") - } - - fn sanitize_package_name(&self, package: &str) -> String { - const ALLOWED_SYMBOLS: &[char] = &['-', '_', '.', '=', '"', ':', '/', '@']; - let sanitized = package - .chars() - .filter(|c| c.is_alphanumeric() || ALLOWED_SYMBOLS.contains(c)) - .collect::(); - - sanitized - } -} - -/// Capture stdout and stderr from the attached container -async fn capture_container_output( - mut attached_container: AttachContainerResults, - stdout: &mut String, - stderr: &mut String, - tx: &SenderWithLogging, -) { - while let Some(output_result) = attached_container.output.next().await { - match output_result { - Ok(output) => match output { - LogOutput::StdOut { message } => { - let message_str = String::from_utf8_lossy(&message).into_owned(); - stdout.push_str(&format!("{message_str}\n")); - tx.send(ToolLog::Result(message_str)).await.ok(); - } - LogOutput::StdErr { message } => { - let message_str = String::from_utf8_lossy(&message).into_owned(); - stderr.push_str(&format!("{message_str}\n")); - tx.send(ToolLog::Result(message_str)).await.ok(); - } - _ => {} - }, - Err(e) => { - let _ = tx.send(ToolLog::Error(e.to_string())).await; - } - } - } -} - -async fn docker_cleanup(docker: &Docker, container_name: &str, image_tag: &str) { - let _ = docker - .stop_container(container_name, None::) - .await; - let _ = tokio::join!( - docker.remove_container( - container_name, - Some(RemoveContainerOptionsBuilder::new().force(true).build()), - ), - docker.remove_image( - image_tag, - Some(RemoveImageOptionsBuilder::new().force(true).build()), - None, - ) - ); -} - -async fn send_log(tx: &SenderWithLogging, message: String) { - let _ = tx.send(ToolLog::Log(message)).await; -} -async fn send_debug(tx: &SenderWithLogging, message: String) { - let _ = tx.send(ToolLog::Debug(message)).await; -} -async fn send_error(tx: &SenderWithLogging, message: String) { - let _ = tx.send(ToolLog::Error(message)).await; -} diff --git a/server/src/tools/system/code_runner/dockerfiles.rs b/server/src/tools/system/code_runner/dockerfiles.rs deleted file mode 100644 index 55284ad..0000000 --- a/server/src/tools/system/code_runner/dockerfiles.rs +++ /dev/null @@ -1,162 +0,0 @@ -use const_format::formatcp; - -use super::CodeLanguage; - -pub fn get_dockerfile(language: &CodeLanguage) -> &'static str { - match language { - CodeLanguage::JavaScript => JS_DOCKERFILE, - CodeLanguage::TypeScript => TS_DOCKERFILE, - CodeLanguage::Python => PYTHON_DOCKERFILE, - CodeLanguage::Rust => RUST_DOCKERFILE, - CodeLanguage::Go => GO_DOCKERFILE, - CodeLanguage::Bash => BASH_DOCKERFILE, - } -} - -pub fn get_dockerfile_info(language: &CodeLanguage) -> (&'static str, &'static str, &'static str) { - let (base_image, file_name, cmd) = match language { - CodeLanguage::JavaScript => (JS_IMAGE, "main.js", "node main.js"), - CodeLanguage::TypeScript => (JS_IMAGE, "main.ts", "pnpm tsx main.ts"), - CodeLanguage::Python => (PYTHON_IMAGE, "main.py", "python main.py"), - CodeLanguage::Rust => (RUST_IMAGE, "main.rs", "./target/debug/temp"), - CodeLanguage::Go => (GO_IMAGE, "main.go", "./temp"), - CodeLanguage::Bash => (BASH_IMAGE, "script.sh", "bash script.sh"), - }; - (base_image, file_name, cmd) -} - -const JS_IMAGE: &str = "node:20-slim"; -const PYTHON_IMAGE: &str = "python:3.13-slim"; -const RUST_IMAGE: &str = "rust:1.85-slim"; -const GO_IMAGE: &str = "golang:1.24"; -const BASH_IMAGE: &str = "bash:5.3"; - -const SET_USER_AND_HOME_DIR: &str = r#" -RUN mkdir -p /app && chown 1000:1000 /app -USER 1000:1000 -RUN mkdir -p /tmp/home -WORKDIR /app -"#; - -const JS_DOCKERFILE: &str = formatcp!( - r#" -FROM {JS_IMAGE} - -ARG DEPENDENCIES -ENV PNPM_HOME="/opt/pnpm" -ENV PATH="$PNPM_HOME:$PATH" - -RUN mkdir -p /opt/pnpm && chown 1000:1000 /opt/pnpm -RUN npm install -g pnpm@9 - -{SET_USER_AND_HOME_DIR} - -RUN pnpm init -RUN if [ -n "$DEPENDENCIES" ]; then pnpm install $DEPENDENCIES; fi - -COPY main.js . - -CMD ["node", "main.js"] -"# -); - -const TS_DOCKERFILE: &str = formatcp!( - r#" -FROM {JS_IMAGE} - -ARG DEPENDENCIES -ENV PNPM_HOME="/opt/pnpm" -ENV PATH="$PNPM_HOME:$PATH" - -RUN mkdir -p /opt/pnpm && chown 1000:1000 /opt/pnpm -RUN npm install -g pnpm@9 - -{SET_USER_AND_HOME_DIR} - -RUN pnpm init -RUN pnpm install tsx $DEPENDENCIES - -COPY main.ts . - -CMD ["pnpm", "tsx", "main.ts"] -"# -); - -const PYTHON_DOCKERFILE: &str = formatcp!( - r#" -FROM {PYTHON_IMAGE} - -ARG DEPENDENCIES -ENV PYTHONUNBUFFERED=1 -ENV PYTHONUSERBASE="/opt/python" -ENV PATH="/opt/python/bin:$PATH" - -RUN mkdir -p /opt/python && chown 1000:1000 /opt/python - -{SET_USER_AND_HOME_DIR} - -RUN if [ -n "$DEPENDENCIES" ]; then pip install --user --no-cache-dir $DEPENDENCIES; fi - -COPY main.py . - -CMD ["python", "main.py"] -"# -); - -const RUST_DOCKERFILE: &str = formatcp!( - r#" -FROM {RUST_IMAGE} -RUN apt-get update -qq && apt-get install -y -qq pkg-config libssl-dev ca-certificates && apt-get clean - -ARG DEPENDENCIES - -{SET_USER_AND_HOME_DIR} - -RUN cargo init --name temp -RUN if [ -n "$DEPENDENCIES" ]; then cargo add $DEPENDENCIES; fi -RUN cargo build - -COPY --chown=1000:1000 main.rs src/ -RUN touch src/main.rs -RUN cargo build - -CMD ["./target/debug/temp"] -"# -); - -const GO_DOCKERFILE: &str = formatcp!( - r#" -FROM {GO_IMAGE} - -ARG DEPENDENCIES - -ENV GOTMPDIR=/opt/gotmpdir GOCACHE=/opt/gocache -RUN mkdir -p /opt/gotmpdir && chown 1000:1000 /opt/gotmpdir -RUN mkdir -p /opt/gocache && chown 1000:1000 /opt/gocache - -{SET_USER_AND_HOME_DIR} - -RUN go mod init temp -RUN if [ -n "$DEPENDENCIES" ]; then go get $DEPENDENCIES; fi - -COPY main.go . -RUN go build - -CMD ["./temp"] -"# -); - -const BASH_DOCKERFILE: &str = formatcp!( - r#" -FROM {BASH_IMAGE} - -ARG DEPENDENCIES -RUN if [ -n "$DEPENDENCIES" ]; then apk add --no-cache $DEPENDENCIES; fi - -{SET_USER_AND_HOME_DIR} - -COPY script.sh . - -CMD ["bash", "script.sh"] -"# -); diff --git a/server/src/tools/system/code_runner/tinirun.rs b/server/src/tools/system/code_runner/tinirun.rs new file mode 100644 index 0000000..571393d --- /dev/null +++ b/server/src/tools/system/code_runner/tinirun.rs @@ -0,0 +1,129 @@ +use rocket::futures::StreamExt; +use tinirun_client::{ + models::{CodeRunnerChunk, CodeRunnerError, CodeRunnerLanguage}, + TinirunClient, +}; + +use crate::{ + tools::{ + core::{ToolLog, ToolResult}, + system::code_runner::CodeLanguage, + ToolError, + }, + utils::SenderWithLogging, +}; + +pub struct TinirunExecutor<'a> { + client: &'a TinirunClient, + lang: CodeLanguage, + timeout_seconds: u32, + memory_limit_mb: u32, + cpu_limit: f32, +} + +#[derive(Debug, Default)] +pub struct TinirunExecutorOptions { + pub timeout_seconds: u32, + pub memory_limit_mb: u32, + pub cpu_limit: f32, +} + +impl<'a> TinirunExecutor<'a> { + pub fn new( + client: &'a TinirunClient, + lang: CodeLanguage, + options: TinirunExecutorOptions, + ) -> Self { + TinirunExecutor { + client, + lang, + timeout_seconds: options.timeout_seconds, + memory_limit_mb: options.memory_limit_mb, + cpu_limit: options.cpu_limit, + } + } + + pub async fn execute( + &self, + code: &str, + dependencies: &[String], + tx: &SenderWithLogging, + ) -> ToolResult { + let input = tinirun_client::models::CodeRunnerInput { + code: code.to_owned(), + lang: self.lang.into(), + dependencies: Some(dependencies.to_vec()), + files: None, + timeout: self.timeout_seconds, + mem_limit_mb: self.memory_limit_mb, + cpu_limit: self.cpu_limit, + }; + let mut stream = match self.client.run_code(&input).await { + Ok(stream) => stream, + Err(err) => { + return Err(ToolError::ToolExecutionError(err.to_string())); + } + }; + while let Some(item) = stream.next().await { + match item { + Ok(event) => match event { + CodeRunnerChunk::Info(log) => tx.send(ToolLog::Log(log)).await.ok(), + CodeRunnerChunk::Debug(log) => tx.send(ToolLog::Debug(log)).await.ok(), + CodeRunnerChunk::Stdout(stdout) => tx.send(ToolLog::Result(stdout)).await.ok(), + CodeRunnerChunk::Stderr(stderr) => tx.send(ToolLog::Result(stderr)).await.ok(), + CodeRunnerChunk::Error(err) => { + tx.send(ToolLog::Error(err.to_string())).await.ok(); + let error_message = match err { + CodeRunnerError::BuildFailed { message, logs } => { + format!("{}\n## Build logs\n{}", message, logs) + } + _ => err.to_string(), + }; + return Err(ToolError::ToolExecutionError(error_message)); + } + CodeRunnerChunk::Result { + stdout, + stderr, + exit_code, + timeout, + } => { + let mut markdown = String::new(); + if let Some(code) = exit_code { + if code != 0 { + markdown += &format!("⚠️ Program exited with code {code}\n\n"); + } + } else if timeout { + markdown += + &format!("⚠️ Timed out after {} seconds\n\n", self.timeout_seconds); + } + + if !stdout.is_empty() { + markdown += &format!("Output:\n{stdout}\n\n"); + } + if !stderr.is_empty() { + markdown += &format!("Stderr:\n{stderr}\n"); + } + + return Ok(markdown); + } + }, + Err(err) => tx.send(ToolLog::Error(err.to_string())).await.ok(), + }; + } + + Err(ToolError::ToolExecutionError("No output".to_string())) + } +} + +impl From for CodeRunnerLanguage { + fn from(language: CodeLanguage) -> Self { + match language { + CodeLanguage::Python => CodeRunnerLanguage::Python, + CodeLanguage::JavaScript => CodeRunnerLanguage::JavaScript, + CodeLanguage::TypeScript => CodeRunnerLanguage::TypeScript, + CodeLanguage::Rust => CodeRunnerLanguage::Rust, + CodeLanguage::Go => CodeRunnerLanguage::Go, + CodeLanguage::Bash => CodeRunnerLanguage::Bash, + } + } +} diff --git a/web/src/components/chat/ChatStreamingToolCalls.tsx b/web/src/components/chat/ChatStreamingToolCalls.tsx index 4bd8336..a79d697 100644 --- a/web/src/components/chat/ChatStreamingToolCalls.tsx +++ b/web/src/components/chat/ChatStreamingToolCalls.tsx @@ -1,5 +1,5 @@ import { ChevronDown, ChevronUp, Loader2, Wrench, X } from "lucide-react"; -import { useMemo, useState } from "react"; +import { useMemo, useRef, useState } from "react"; import Markdown from "react-markdown"; import { getToolIcon, getToolTypeLabel } from "@/components/ToolsManager"; @@ -201,11 +201,14 @@ function StreamingToolCall({ } function DebugLogsContent({ children }: { children: React.ReactNode }) { - const { scrollRef } = useAutoScroll(); + const contentRef = useRef(null); + const { scrollRef } = useAutoScroll({ contentRef }); return ( -
- {children} +
+
+ {children} +
); } diff --git a/web/src/components/chat/messages/ChatMessageToolLogs.tsx b/web/src/components/chat/messages/ChatMessageToolLogs.tsx index 46a527a..04e6873 100644 --- a/web/src/components/chat/messages/ChatMessageToolLogs.tsx +++ b/web/src/components/chat/messages/ChatMessageToolLogs.tsx @@ -1,5 +1,5 @@ import { ChevronDown, ChevronUp } from "lucide-react"; -import { useState } from "react"; +import { useRef, useState } from "react"; import { Button } from "@/components/ui/button"; import { useAutoScroll } from "@/components/ui/chat/hooks/useAutoScroll"; @@ -18,6 +18,9 @@ export default function ChatMessageToolLogs({ }) { const [showLogs, setShowLogs] = useState(initialOpen ?? false); + const contentRef = useRef(null); + const { scrollRef } = useAutoScroll({ contentRef }); + return ( @@ -28,27 +31,19 @@ export default function ChatMessageToolLogs({ - - {logs.map((log, index) => ( -
- {log} -
- ))} -
+
+
+ {logs.map((log, index) => ( +
+ {log} +
+ ))} +
+
); } - -function LogsContent({ children }: { children: React.ReactNode }) { - const { scrollRef } = useAutoScroll(); - - return ( -
- {children} -
- ); -} diff --git a/web/src/components/chat/messages/ChatMessageToolResult.tsx b/web/src/components/chat/messages/ChatMessageToolResult.tsx index 8a7071a..f511d74 100644 --- a/web/src/components/chat/messages/ChatMessageToolResult.tsx +++ b/web/src/components/chat/messages/ChatMessageToolResult.tsx @@ -23,7 +23,7 @@ export default function ChatMessageToolResult({ message: components["schemas"]["ChatRsMessage"]; tools?: components["schemas"]["GetAllToolsResponse"]; }) { - const [showOutput, setShowOutput] = useState(false); + const [showOutput, setShowOutput] = useState(true); const tool = useMemo(() => { if (!message.meta.tool_call) return null; diff --git a/web/src/components/ui/chat/chat-message-list.tsx b/web/src/components/ui/chat/chat-message-list.tsx index ddc1813..1e50875 100644 --- a/web/src/components/ui/chat/chat-message-list.tsx +++ b/web/src/components/ui/chat/chat-message-list.tsx @@ -14,7 +14,6 @@ const ChatMessageList = React.forwardRef( const { scrollRef, isAtBottom, scrollToBottom, disableAutoScroll } = useAutoScroll({ smooth, - contentRef, }); From 40049f9206caa69eb75da5382d94dd1f377eec9a Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 22 Jun 2026 02:16:52 -0400 Subject: [PATCH 021/111] add new server --- server-new/.env.example | 10 + server-new/.gitignore | 20 + server-new/Cargo.lock | 1588 ++++++++++++++++++++++++++ server-new/Cargo.toml | 43 + server-new/Dockerfile | 45 + server-new/README.md | 134 +++ server-new/src/config.rs | 78 ++ server-new/src/error.rs | 88 ++ server-new/src/extractors/mod.rs | 1 + server-new/src/extractors/session.rs | 70 ++ server-new/src/lib.rs | 23 + server-new/src/main.rs | 96 ++ server-new/src/plugins/logging.rs | 49 + server-new/src/plugins/mod.rs | 3 + server-new/src/plugins/security.rs | 32 + server-new/src/plugins/session.rs | 33 + server-new/src/routes/auth.rs | 36 + server-new/src/routes/health.rs | 9 + server-new/src/routes/hello.rs | 15 + server-new/src/routes/mod.rs | 19 + server-new/src/state.rs | 33 + 21 files changed, 2425 insertions(+) create mode 100644 server-new/.env.example create mode 100644 server-new/.gitignore create mode 100644 server-new/Cargo.lock create mode 100644 server-new/Cargo.toml create mode 100644 server-new/Dockerfile create mode 100644 server-new/README.md create mode 100644 server-new/src/config.rs create mode 100644 server-new/src/error.rs create mode 100644 server-new/src/extractors/mod.rs create mode 100644 server-new/src/extractors/session.rs create mode 100644 server-new/src/lib.rs create mode 100644 server-new/src/main.rs create mode 100644 server-new/src/plugins/logging.rs create mode 100644 server-new/src/plugins/mod.rs create mode 100644 server-new/src/plugins/security.rs create mode 100644 server-new/src/plugins/session.rs create mode 100644 server-new/src/routes/auth.rs create mode 100644 server-new/src/routes/health.rs create mode 100644 server-new/src/routes/hello.rs create mode 100644 server-new/src/routes/mod.rs create mode 100644 server-new/src/state.rs diff --git a/server-new/.env.example b/server-new/.env.example new file mode 100644 index 0000000..c83e520 --- /dev/null +++ b/server-new/.env.example @@ -0,0 +1,10 @@ +# Server Configuration +RS_CHAT_HOST=127.0.0.1 +RS_CHAT_PORT=8080 + +# Auth +RS_CHAT_COOKIE_KEY= # hex secret >=32 bytes, e.g. `openssl rand --hex 32` + +# Logging +RS_CHAT_LOG_LEVEL=info +RS_CHAT_REQUEST_ID_HEADER=x-request-id diff --git a/server-new/.gitignore b/server-new/.gitignore new file mode 100644 index 0000000..b4253ee --- /dev/null +++ b/server-new/.gitignore @@ -0,0 +1,20 @@ +# Rust +/target + +# Env files +.env* +!.env.example + +# OS files +.DS_Store + +# Build artifacts +*.exe +*.dll +*.so +*.dylib + +# Temporary files +tmp/ +temp/ +*.tmp diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock new file mode 100644 index 0000000..c82280c --- /dev/null +++ b/server-new/Cargo.lock @@ -0,0 +1,1588 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "aead" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0" +dependencies = [ + "crypto-common", + "generic-array", +] + +[[package]] +name = "aes" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0" +dependencies = [ + "cfg-if", + "cipher", + "cpufeatures", +] + +[[package]] +name = "aes-gcm" +version = "0.10.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "831010a0f742e1209b3bcea8fab6a8e149051ba6099432c8cb2cc117dec3ead1" +dependencies = [ + "aead", + "aes", + "cipher", + "ctr", + "ghash", + "subtle", +] + +[[package]] +name = "aho-corasick" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ddd31a130427c27518df266943a5308ed92d4b226cc639f5a8f1002816174301" +dependencies = [ + "memchr", +] + +[[package]] +name = "anyhow" +version = "1.0.102" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" + +[[package]] +name = "async-trait" +version = "0.1.89" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9035ad2d096bed7955a320ee7e2230574d28fd3c3a0f186cbea1ff3c7eed5dbb" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "atomic" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a89cbf775b137e9b968e67227ef7f775587cde3fd31b0d8599dbd0f598a48340" +dependencies = [ + "bytemuck", +] + +[[package]] +name = "atomic-waker" +version = "1.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0" + +[[package]] +name = "axum" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "31b698c5f9a010f6573133b09e0de5408834d0c82f8d7475a89fc1867a71cd90" +dependencies = [ + "axum-core", + "bytes", + "form_urlencoded", + "futures-util", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-util", + "itoa", + "matchit", + "memchr", + "mime", + "percent-encoding", + "pin-project-lite", + "serde_core", + "serde_json", + "serde_path_to_error", + "serde_urlencoded", + "sync_wrapper", + "tokio", + "tower", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "axum-core" +version = "0.5.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08c78f31d7b1291f7ee735c1c6780ccde7785daae9a9206026862dab7d8792d1" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "http-body-util", + "mime", + "pin-project-lite", + "sync_wrapper", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "axum-helmet" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4233d7fef77c993a0251c2e4610cc6e4e02c01bbd669f8f323ec9e8376d947c7" +dependencies = [ + "helmet-core", + "http", + "pin-project-lite", + "tower", + "tower-service", +] + +[[package]] +name = "axum-plugin" +version = "0.2.0" +source = "git+https://git.fasharp.io/fa-sharp/axum-plugin?rev=9f72278b3c#9f72278b3c8ae57897ac04af54a6643bb9e5a25c" +dependencies = [ + "anyhow", + "axum", + "axum-plugin-macros", + "futures", + "type-map", +] + +[[package]] +name = "axum-plugin-macros" +version = "0.1.0" +source = "git+https://git.fasharp.io/fa-sharp/axum-plugin?rev=9f72278b3c#9f72278b3c8ae57897ac04af54a6643bb9e5a25c" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "base64" +version = "0.22.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" + +[[package]] +name = "bitflags" +version = "2.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4388bee8683e3d04af747c73422af53102d2bd24d9eadb6cbc100baef4b43f8" + +[[package]] +name = "block-buffer" +version = "0.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" +dependencies = [ + "generic-array", +] + +[[package]] +name = "bumpalo" +version = "3.20.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649" + +[[package]] +name = "bytemuck" +version = "1.25.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8efb64bd706a16a1bdde310ae86b351e4d21550d98d056f22f8a7f7a2183fec" + +[[package]] +name = "bytes" +version = "1.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ae3f5d315924270530207e2a68396c3cc547f6dca3fbdca317cfb1a51edb593" + +[[package]] +name = "cfg-if" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" + +[[package]] +name = "cipher" +version = "0.4.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" +dependencies = [ + "crypto-common", + "inout", +] + +[[package]] +name = "cookie" +version = "0.18.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4ddef33a339a91ea89fb53151bd0a4689cfce27055c291dfa69945475d22c747" +dependencies = [ + "aes-gcm", + "base64", + "hkdf", + "hmac", + "percent-encoding", + "rand 0.8.6", + "sha2", + "subtle", + "time", + "version_check", +] + +[[package]] +name = "cpufeatures" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" +dependencies = [ + "libc", +] + +[[package]] +name = "crossbeam-channel" +version = "0.5.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "82b8f8f868b36967f9606790d1903570de9ceaf870a7bf9fbbd3016d636a2cb2" +dependencies = [ + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-utils" +version = "0.8.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" + +[[package]] +name = "crypto-common" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" +dependencies = [ + "generic-array", + "rand_core 0.6.4", + "typenum", +] + +[[package]] +name = "ctr" +version = "0.9.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0369ee1ad671834580515889b80f2ea915f23b8be8d0daa4bbaf2ac5c7590835" +dependencies = [ + "cipher", +] + +[[package]] +name = "deranged" +version = "0.5.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7cd812cc2bc1d69d4764bd80df88b4317eaef9e773c75226407d9bc0876b211c" +dependencies = [ + "serde_core", +] + +[[package]] +name = "digest" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" +dependencies = [ + "block-buffer", + "crypto-common", + "subtle", +] + +[[package]] +name = "dotenvy" +version = "0.15.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1aaf95b3e5c8f23aa320147307562d361db0ae0d51242340f558153b4eb2439b" + +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys", +] + +[[package]] +name = "figment" +version = "0.10.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8cb01cd46b0cf372153850f4c6c272d9cbea2da513e07538405148f95bd789f3" +dependencies = [ + "atomic", + "pear", + "serde", + "uncased", + "version_check", +] + +[[package]] +name = "form_urlencoded" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb4cb245038516f5f85277875cdaa4f7d2c9a0fa0468de06ed190163b1581fcf" +dependencies = [ + "percent-encoding", +] + +[[package]] +name = "futures" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b147ee9d1f6d097cef9ce628cd2ee62288d963e16fb287bd9286455b241382d" +dependencies = [ + "futures-channel", + "futures-core", + "futures-executor", + "futures-io", + "futures-sink", + "futures-task", + "futures-util", +] + +[[package]] +name = "futures-channel" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "07bbe89c50d7a535e539b8c17bc0b49bdb77747034daa8087407d655f3f7cc1d" +dependencies = [ + "futures-core", + "futures-sink", +] + +[[package]] +name = "futures-core" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d" + +[[package]] +name = "futures-executor" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "baf29c38818342a3b26b5b923639e7b1f4a61fc5e76102d4b1981c6dc7a7579d" +dependencies = [ + "futures-core", + "futures-task", + "futures-util", +] + +[[package]] +name = "futures-io" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cecba35d7ad927e23624b22ad55235f2239cfa44fd10428eecbeba6d6a717718" + +[[package]] +name = "futures-macro" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e835b70203e41293343137df5c0664546da5745f82ec9b84d40be8336958447b" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "futures-sink" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c39754e157331b013978ec91992bde1ac089843443c49cbc7f46150b0fad0893" + +[[package]] +name = "futures-task" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "037711b3d59c33004d3856fbdc83b99d4ff37a24768fa1be9ce3538a1cde4393" + +[[package]] +name = "futures-util" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" +dependencies = [ + "futures-channel", + "futures-core", + "futures-io", + "futures-macro", + "futures-sink", + "futures-task", + "memchr", + "pin-project-lite", + "slab", +] + +[[package]] +name = "generic-array" +version = "0.14.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +dependencies = [ + "typenum", + "version_check", +] + +[[package]] +name = "getrandom" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0" +dependencies = [ + "cfg-if", + "libc", + "wasi", +] + +[[package]] +name = "getrandom" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" +dependencies = [ + "cfg-if", + "libc", + "r-efi 5.3.0", + "wasip2", +] + +[[package]] +name = "getrandom" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "300e883d756b2e4ec94e02791f39b04b522276138852cfc41d9fb7e904106099" +dependencies = [ + "cfg-if", + "libc", + "r-efi 6.0.0", +] + +[[package]] +name = "ghash" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0d8a4362ccb29cb0b265253fb0a2728f592895ee6854fd9bc13f2ffda266ff1" +dependencies = [ + "opaque-debug", + "polyval", +] + +[[package]] +name = "helmet-core" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a865b9c8b67316ab132710af828252e764cdf2195bbcd72a23b96127150d9de" + +[[package]] +name = "hex" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" + +[[package]] +name = "hkdf" +version = "0.12.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b5f8eb2ad728638ea2c7d47a21db23b7b58a72ed6a38256b8a1849f15fbbdf7" +dependencies = [ + "hmac", +] + +[[package]] +name = "hmac" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" +dependencies = [ + "digest", +] + +[[package]] +name = "http" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6970f50e31d6fc17d3fa27329444bfa74e196cf62e95052a3f6fee181dba6425" +dependencies = [ + "bytes", + "itoa", +] + +[[package]] +name = "http-body" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1efedce1fb8e6913f23e0c92de8e62cd5b772a67e7b3946df930a62566c93184" +dependencies = [ + "bytes", + "http", +] + +[[package]] +name = "http-body-util" +version = "0.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b021d93e26becf5dc7e1b75b1bed1fd93124b374ceb73f43d4d4eafec896a64a" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "pin-project-lite", +] + +[[package]] +name = "httparse" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87" + +[[package]] +name = "httpdate" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" + +[[package]] +name = "hyper" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "55281c53a1894c864990125767da440a4e630446785086f52523b20033b74498" +dependencies = [ + "atomic-waker", + "bytes", + "futures-channel", + "futures-core", + "http", + "http-body", + "httparse", + "httpdate", + "itoa", + "pin-project-lite", + "smallvec", + "tokio", +] + +[[package]] +name = "hyper-util" +version = "0.1.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0" +dependencies = [ + "bytes", + "http", + "http-body", + "hyper", + "pin-project-lite", + "tokio", + "tower-service", +] + +[[package]] +name = "inlinable_string" +version = "0.1.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8fae54786f62fb2918dcfae3d568594e50eb9b5c25bf04371af6fe7516452fb" + +[[package]] +name = "inout" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "879f10e63c20629ecabbb64a8010319738c66a5cd0c29b02d63d272b03751d01" +dependencies = [ + "generic-array", +] + +[[package]] +name = "itoa" +version = "1.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" + +[[package]] +name = "js-sys" +version = "0.3.102" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "03d04c30968dffe80775bd4d7fb676131cd04a1fb46d2686dbffbaec2d9dfd31" +dependencies = [ + "cfg-if", + "futures-util", + "wasm-bindgen", +] + +[[package]] +name = "lazy_static" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" + +[[package]] +name = "libc" +version = "0.2.186" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" + +[[package]] +name = "lock_api" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "224399e74b87b5f3557511d98dff8b14089b3dadafcab6bb93eab67d3aace965" +dependencies = [ + "scopeguard", +] + +[[package]] +name = "log" +version = "0.4.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ceec5bc11778974d1bcb055b18002eba7f4b3518b6a0081b3af5f21666da9ad" + +[[package]] +name = "matchers" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d1525a2a28c7f4fa0fc98bb91ae755d1e2d1505079e05539e35bc876b5d65ae9" +dependencies = [ + "regex-automata", +] + +[[package]] +name = "matchit" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" + +[[package]] +name = "memchr" +version = "2.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "88904434abc2901f197fe8cc55f0445e7ded921dba5911dad2e2b39b48e663c4" + +[[package]] +name = "mime" +version = "0.3.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a" + +[[package]] +name = "mio" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "02bd0af71c67b473010cbbc60715ee815645a4dc942899111f494b4b737d6fda" +dependencies = [ + "libc", + "wasi", + "windows-sys", +] + +[[package]] +name = "nu-ansi-term" +version = "0.50.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" +dependencies = [ + "windows-sys", +] + +[[package]] +name = "num-conv" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521739c6d2bac4aa25192232afe6841231376b2b26d4d9fae5ecf8ca5772e441" + +[[package]] +name = "once_cell" +version = "1.21.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" + +[[package]] +name = "opaque-debug" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" + +[[package]] +name = "parking_lot" +version = "0.12.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93857453250e3077bd71ff98b6a65ea6621a19bb0f559a85248955ac12c45a1a" +dependencies = [ + "lock_api", + "parking_lot_core", +] + +[[package]] +name = "parking_lot_core" +version = "0.9.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2621685985a2ebf1c516881c026032ac7deafcda1a2c9b7850dc81e3dfcb64c1" +dependencies = [ + "cfg-if", + "libc", + "redox_syscall", + "smallvec", + "windows-link", +] + +[[package]] +name = "pear" +version = "0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bdeeaa00ce488657faba8ebf44ab9361f9365a97bd39ffb8a60663f57ff4b467" +dependencies = [ + "inlinable_string", + "pear_codegen", + "yansi", +] + +[[package]] +name = "pear_codegen" +version = "0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4bab5b985dc082b345f812b7df84e1bef27e7207b39e448439ba8bd69c93f147" +dependencies = [ + "proc-macro2", + "proc-macro2-diagnostics", + "quote", + "syn", +] + +[[package]] +name = "percent-encoding" +version = "2.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" + +[[package]] +name = "pin-project-lite" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" + +[[package]] +name = "polyval" +version = "0.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25" +dependencies = [ + "cfg-if", + "cpufeatures", + "opaque-debug", + "universal-hash", +] + +[[package]] +name = "powerfmt" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "439ee305def115ba05938db6eb1644ff94165c5ab5e9420d1c1bcedbba909391" + +[[package]] +name = "ppv-lite86" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85eae3c4ed2f50dcfe72643da4befc30deadb458a9b590d720cde2f2b1e97da9" +dependencies = [ + "zerocopy", +] + +[[package]] +name = "proc-macro2" +version = "1.0.106" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "proc-macro2-diagnostics" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "af066a9c399a26e020ada66a034357a868728e72cd426f3adcd35f80d88d88c8" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "version_check", + "yansi", +] + +[[package]] +name = "quote" +version = "1.0.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dfbc457d0c7a0759a614551b11a6409e5951f6c7537be1f1b7682b9ae9230368" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "r-efi" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" + +[[package]] +name = "r-efi" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" + +[[package]] +name = "rand" +version = "0.8.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5ca0ecfa931c29007047d1bc58e623ab12e5590e8c7cc53200d5202b69266d8a" +dependencies = [ + "libc", + "rand_chacha 0.3.1", + "rand_core 0.6.4", +] + +[[package]] +name = "rand" +version = "0.9.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "44c5af06bb1b7d3216d91932aed5265164bf384dc89cd6ba05cf59a35f5f76ea" +dependencies = [ + "rand_chacha 0.9.0", + "rand_core 0.9.5", +] + +[[package]] +name = "rand_chacha" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6c10a63a0fa32252be49d21e7709d4d4baf8d231c2dbce1eaa8141b9b127d88" +dependencies = [ + "ppv-lite86", + "rand_core 0.6.4", +] + +[[package]] +name = "rand_chacha" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" +dependencies = [ + "ppv-lite86", + "rand_core 0.9.5", +] + +[[package]] +name = "rand_core" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec0be4795e2f6a28069bec0b5ff3e2ac9bafc99e6a9a7dc3547996c5c816922c" +dependencies = [ + "getrandom 0.2.17", +] + +[[package]] +name = "rand_core" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c" +dependencies = [ + "getrandom 0.3.4", +] + +[[package]] +name = "redox_syscall" +version = "0.5.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d" +dependencies = [ + "bitflags", +] + +[[package]] +name = "regex-automata" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e1dd4122fc1595e8162618945476892eefca7b88c52820e74af6262213cae8f" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax", +] + +[[package]] +name = "regex-syntax" +version = "0.8.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" + +[[package]] +name = "rs-chat-api" +version = "0.1.0" +dependencies = [ + "anyhow", + "axum", + "axum-helmet", + "axum-plugin", + "dotenvy", + "figment", + "hex", + "serde", + "serde_json", + "time", + "tokio", + "tower", + "tower-http", + "tower-sessions", + "tracing", + "tracing-appender", + "tracing-subscriber", +] + +[[package]] +name = "rustc-hash" +version = "2.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94300abf3f1ae2e2b8ffb7b58043de3d399c73fa6f4b73826402a5c457614dbe" + +[[package]] +name = "rustversion" +version = "1.0.22" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d" + +[[package]] +name = "ryu" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" + +[[package]] +name = "scopeguard" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" + +[[package]] +name = "serde" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "serde_json" +version = "1.0.150" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e8014e44b4736ed0538adeecded0fce2a272f22dc9578a7eb6b2d9993c74cfb9" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + +[[package]] +name = "serde_path_to_error" +version = "0.1.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10a9ff822e371bb5403e391ecd83e182e0e77ba7f6fe0160b795797109d1b457" +dependencies = [ + "itoa", + "serde", + "serde_core", +] + +[[package]] +name = "serde_urlencoded" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3491c14715ca2294c4d6a88f15e84739788c1d030eed8c110436aafdaa2f3fd" +dependencies = [ + "form_urlencoded", + "itoa", + "ryu", + "serde", +] + +[[package]] +name = "sha2" +version = "0.10.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" +dependencies = [ + "cfg-if", + "cpufeatures", + "digest", +] + +[[package]] +name = "sharded-slab" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f40ca3c46823713e0d4209592e8d6e826aa57e928f09752619fc696c499637f6" +dependencies = [ + "lazy_static", +] + +[[package]] +name = "signal-hook-registry" +version = "1.4.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4db69cba1110affc0e9f7bcd48bbf87b3f4fc7c61fc9155afd4c469eb3d6c1b" +dependencies = [ + "errno", + "libc", +] + +[[package]] +name = "slab" +version = "0.4.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c790de23124f9ab44544d7ac05d60440adc586479ce501c1d6d7da3cd8c9cf5" + +[[package]] +name = "smallvec" +version = "1.15.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ed6a63f02c8539c91a8685a86f4099661ba3da017932f6ebbea6de3f0fa7c90" + +[[package]] +name = "socket2" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52d1cfed4120b4d927bf7c0f86d2087a4a7d6027c906d9f9d525a80573b9be51" +dependencies = [ + "libc", + "windows-sys", +] + +[[package]] +name = "subtle" +version = "2.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" + +[[package]] +name = "symlink" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a7973cce6668464ea31f176d85b13c7ab3bba2cb3b77a2ed26abd7801688010a" + +[[package]] +name = "syn" +version = "2.0.118" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1b9ae57f904213ebb649ce6895b8a66c66f0203b9319718f69a5612a065b1422" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "sync_wrapper" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" + +[[package]] +name = "thiserror" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" +dependencies = [ + "thiserror-impl", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "thread_local" +version = "1.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f60246a4944f24f6e018aa17cdeffb7818b76356965d03b07d6a9886e8962185" +dependencies = [ + "cfg-if", +] + +[[package]] +name = "time" +version = "0.3.49" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "711a53c2d47bbd818258c498c8dbfe186a2526c631495cfe7e078567f86b8469" +dependencies = [ + "deranged", + "num-conv", + "powerfmt", + "serde_core", + "time-core", + "time-macros", +] + +[[package]] +name = "time-core" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9e1c906769ad99c88eaa54e728060edef082f8e358ff32030cb7c7d315e81109" + +[[package]] +name = "time-macros" +version = "0.2.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "71c652a3727a9cbb9a02f707f530b618ce00d0ccd762009c8c23bd191df3c17d" +dependencies = [ + "num-conv", + "time-core", +] + +[[package]] +name = "tokio" +version = "1.52.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe" +dependencies = [ + "libc", + "mio", + "pin-project-lite", + "signal-hook-registry", + "socket2", + "tokio-macros", + "windows-sys", +] + +[[package]] +name = "tokio-macros" +version = "2.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tower" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebe5ef63511595f1344e2d5cfa636d973292adc0eec1f0ad45fae9f0851ab1d4" +dependencies = [ + "futures-core", + "futures-util", + "pin-project-lite", + "sync_wrapper", + "tokio", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "tower-cookies" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "151b5a3e3c45df17466454bb74e9ecedecc955269bdedbf4d150dfa393b55a36" +dependencies = [ + "axum-core", + "cookie", + "futures-util", + "http", + "parking_lot", + "pin-project-lite", + "tower-layer", + "tower-service", +] + +[[package]] +name = "tower-http" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b11f75e912b0c2be01b63d8cf8057b8c3f97cf34abb3d431a3a4c8675498e233" +dependencies = [ + "bitflags", + "bytes", + "http", + "http-body", + "http-body-util", + "percent-encoding", + "pin-project-lite", + "tokio", + "tower-layer", + "tower-service", + "tracing", + "uuid", +] + +[[package]] +name = "tower-layer" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "121c2a6cda46980bb0fcd1647ffaf6cd3fc79a013de288782836f6df9c48780e" + +[[package]] +name = "tower-service" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8df9b6e13f2d32c91b9bd719c00d1958837bc7dec474d94952798cc8e69eeec3" + +[[package]] +name = "tower-sessions" +version = "0.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "518dca34b74a17cadfcee06e616a09d2bd0c3984eff1769e1e76d58df978fc78" +dependencies = [ + "async-trait", + "http", + "time", + "tokio", + "tower-cookies", + "tower-layer", + "tower-service", + "tower-sessions-core", + "tower-sessions-memory-store", + "tracing", +] + +[[package]] +name = "tower-sessions-core" +version = "0.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "568531ec3dfcf3ffe493de1958ae5662a0284ac5d767476ecdb6a34ff8c6b06c" +dependencies = [ + "async-trait", + "axum-core", + "base64", + "futures", + "http", + "parking_lot", + "rand 0.9.4", + "serde", + "serde_json", + "thiserror", + "time", + "tokio", + "tracing", +] + +[[package]] +name = "tower-sessions-memory-store" +version = "0.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "713fabf882b6560a831e2bbed6204048b35bdd60e50bbb722902c74f8df33460" +dependencies = [ + "async-trait", + "time", + "tokio", + "tower-sessions-core", +] + +[[package]] +name = "tracing" +version = "0.1.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" +dependencies = [ + "log", + "pin-project-lite", + "tracing-attributes", + "tracing-core", +] + +[[package]] +name = "tracing-appender" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "050686193eb999b4bb3bc2acfa891a13da00f79734704c4b8b4ef1a10b368a3c" +dependencies = [ + "crossbeam-channel", + "symlink", + "thiserror", + "time", + "tracing-subscriber", +] + +[[package]] +name = "tracing-attributes" +version = "0.1.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7490cfa5ec963746568740651ac6781f701c9c5ea257c58e057f3ba8cf69e8da" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tracing-core" +version = "0.1.36" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" +dependencies = [ + "once_cell", + "valuable", +] + +[[package]] +name = "tracing-log" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ee855f1f400bd0e5c02d150ae5de3840039a3f54b025156404e34c23c03f47c3" +dependencies = [ + "log", + "once_cell", + "tracing-core", +] + +[[package]] +name = "tracing-serde" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "704b1aeb7be0d0a84fc9828cae51dab5970fee5088f83d1dd7ee6f6246fc6ff1" +dependencies = [ + "serde", + "tracing-core", +] + +[[package]] +name = "tracing-subscriber" +version = "0.3.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb7f578e5945fb242538965c2d0b04418d38ec25c79d160cd279bf0731c8d319" +dependencies = [ + "matchers", + "nu-ansi-term", + "once_cell", + "regex-automata", + "serde", + "serde_json", + "sharded-slab", + "smallvec", + "thread_local", + "tracing", + "tracing-core", + "tracing-log", + "tracing-serde", +] + +[[package]] +name = "type-map" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb30dbbd9036155e74adad6812e9898d03ec374946234fbcebd5dfc7b9187b90" +dependencies = [ + "rustc-hash", +] + +[[package]] +name = "typenum" +version = "1.20.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20" + +[[package]] +name = "uncased" +version = "0.9.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e1b88fcfe09e89d3866a5c11019378088af2d24c3fbd4f0543f96b479ec90697" +dependencies = [ + "version_check", +] + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "universal-hash" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea" +dependencies = [ + "crypto-common", + "subtle", +] + +[[package]] +name = "uuid" +version = "1.23.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "144d6b123cef80b301b8f72a9e2ca4370ddec21950d0a103dd22c437006d2db7" +dependencies = [ + "getrandom 0.4.3", + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "valuable" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65" + +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" + +[[package]] +name = "wasi" +version = "0.11.1+wasi-snapshot-preview1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" + +[[package]] +name = "wasip2" +version = "1.0.4+wasi-0.2.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b67efb37e106e55ce722a510d6b5f9c17f083e5fc79afc2badeb12cc313d9487" +dependencies = [ + "wit-bindgen", +] + +[[package]] +name = "wasm-bindgen" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ddb3f79143bced6de84270411622a2699cee572fc0875aeaf1e7867cf9fca1a" +dependencies = [ + "cfg-if", + "once_cell", + "rustversion", + "wasm-bindgen-macro", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-macro" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4e21a184b13fb19e157296e2c46056aec9092264fab83e4ba59e68c61b323c3d" +dependencies = [ + "quote", + "wasm-bindgen-macro-support", +] + +[[package]] +name = "wasm-bindgen-macro-support" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fecefd9c35bd935a20fc3fc344b5f29138961e4f47fb03297d88f2587afb5ebd" +dependencies = [ + "bumpalo", + "proc-macro2", + "quote", + "syn", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-shared" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "23939e44bb9a5d7576fa2b563dc2e136628f1224e88a8deed09e04858b77871f" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "windows-link" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" + +[[package]] +name = "windows-sys" +version = "0.61.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc" +dependencies = [ + "windows-link", +] + +[[package]] +name = "wit-bindgen" +version = "0.57.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" + +[[package]] +name = "yansi" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfe53a6657fd280eaa890a3bc59152892ffa3e30101319d168b781ed6529b049" + +[[package]] +name = "zerocopy" +version = "0.8.52" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ce1022995ff5ff5d841ad7d994facc23098cd40152f2c1d11cd607c6f530653f" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.8.52" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ae7f38b72ec2a254e2b87ef277cf2cd4fb97cbebf944faa6f33354da0867930" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "zmij" +version = "1.0.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml new file mode 100644 index 0000000..e154924 --- /dev/null +++ b/server-new/Cargo.toml @@ -0,0 +1,43 @@ +[package] +name = "rs-chat-api" +version = "0.1.0" +edition = "2024" +description = "LLM chat application" + +[dependencies] +anyhow = "1.0.102" +axum = { version = "0.8.9", features = ["json", "query"] } +axum-helmet = "1.0.2" +axum-plugin = { + git = "https://git.fasharp.io/fa-sharp/axum-plugin", + rev = "9f72278b3c" +} +dotenvy = "0.15.7" +figment = { version = "0.10.19", features = ["env"] } +hex = "0.4.3" +serde = { version = "1.0.228", features = ["derive"] } +serde_json = "1.0.150" +time = { version = "=0.3.49", default-features = false } +tokio = { + version = "1.52.3", + default-features = false, + features = ["macros", "net", "rt", "rt-multi-thread", "signal"] +} +tower = { version = "0.5", default-features = false } +tower-http = { + version = "0.7.0", + features = [ + "limit", + "request-id", + "timeout", + "trace", + ] +} +tower-sessions = { + version = "0.15.0", + default-features = false, + features = ["axum-core", "memory-store", "private"] +} +tracing = "0.1.44" +tracing-appender = "0.2.5" +tracing-subscriber = { version = "0.3.23", features = ["env-filter", "json"] } diff --git a/server-new/Dockerfile b/server-new/Dockerfile new file mode 100644 index 0000000..88b7054 --- /dev/null +++ b/server-new/Dockerfile @@ -0,0 +1,45 @@ +# Image versions (can be overridden by args when building) +ARG RUST_VERSION=1.96 +ARG DEBIAN_VERSION=bookworm + +### Build server ### +FROM rust:${RUST_VERSION}-slim-${DEBIAN_VERSION} AS build +WORKDIR /app + +# Copy all necessary files to build the server +COPY Cargo.lock Cargo.toml ./ +COPY ./src ./src +# COPY ./migrations ./migrations +# etc... + +ARG pkg=rs-chat-api + +RUN --mount=type=cache,id=rust_target,target=/app/target \ + --mount=type=cache,id=cargo_registry,target=/usr/local/cargo/registry \ + --mount=type=cache,id=cargo_git,target=/usr/local/cargo/git \ + set -eux; \ + cargo build --package $pkg --release --locked; \ + objcopy --compress-debug-sections target/release/$pkg ./run-server + + +### Run server ### +FROM debian:${DEBIAN_VERSION}-slim AS run + +# Create non-root user +ARG UID=10001 +RUN adduser \ + --disabled-password \ + --gecos "" \ + --home "/home/appuser" \ + --shell "/sbin/nologin" \ + --uid "${UID}" \ + appuser +USER appuser + +# Copy server binary +COPY --from=build --chown=appuser /app/run-server /usr/local/bin/ + +# Run server +WORKDIR /app +ENV RS_CHAT_HOST=0.0.0.0 +CMD ["run-server"] diff --git a/server-new/README.md b/server-new/README.md new file mode 100644 index 0000000..058045e --- /dev/null +++ b/server-new/README.md @@ -0,0 +1,134 @@ +# Axum Web Service Template + +A production-ready template for building web services with Rust and Axum. + +## Features + +- **Axum** - Fast and ergonomic web framework +- **Configuration Management** - Environment-based config with `figment` +- **Structured API Errors** - JSON error responses with an `AppError` type for route handlers +- **Structured Logging** - JSON logging in production with `tracing` +- **Secure Defaults** - Default HTTP security headers, request body limit and timeout with `tower-http` +- **Optional Request Logging** - Request IDs and HTTP request/response logs with `tower-http` +- **Graceful Shutdown** - Handles SIGTERM and SIGINT signals +- ️**Plugin Architecture** - Modular app initialization with `axum-plugin` +- **Optional OpenAPI** - API documentation with `aide` (optional) +- **Docker / OCI** - Dockerfile with sensible defaults for quick deployment + +## Usage + +### Using cargo-generate + +Install cargo-generate if you haven't already: + +```bash +cargo install cargo-generate +``` + +Generate a new project from this template: + +```bash +cargo generate --git https://git.fasharp.io/fa-sharp/axum-template +``` + +You'll be prompted for: +- **Project name**: The name of your new project +- **Project description**: A brief description +- **Environment variable prefix**: Prefix for env vars (e.g., `APP` for `APP_HOST`, `APP_PORT`) +- **Default port**: The server's default port +- **Default log level**: trace, debug, info, warn, or error +- **Include request logging**: Whether to include request ID and request/response logging middleware +- **Include aide**: Whether to include OpenAPI documentation support + +## Configuration + +Configuration is loaded from environment variables and validated in the `config.rs` file. The variable prefix is configurable during template generation. + +Example with `APP` prefix: + +```bash +# Required +APP_API_KEY=your-secret-key + +# Optional (defaults shown) +APP_HOST=127.0.0.1 +APP_PORT=8080 +APP_LOG_LEVEL=info +APP_REQUEST_ID_HEADER=x-request-id +``` + +In development, you can use the `.env` file to set environment variables. + +## Project Structure + +``` +. +├── src/ +│ ├── routes/ # API routes +│ ├── plugins/ # Axum plugins +│ ├── config.rs # Configuration management +│ ├── error.rs # Structured API error handling +│ ├── lib.rs # Axum server setup +│ ├── main.rs # Entry point +│ └── state.rs # Axum server state +├── Cargo.toml # Dependencies +├── .env # Local environment variables +└── .env.example # Example environment variables +``` + +## Development + +```bash +# Run in development mode (loads .env file) +cargo run + +# Run with custom log level +APP_LOG_LEVEL=debug cargo run + +# Build for production +cargo build --release +``` + +## Adding Routes + +This template uses `axum-plugin` for modular initialization. To add routes: + +1. Create a new plugin in a separate module +2. Register it in `lib.rs`: + +```rust +pub async fn create_app() -> anyhow::Result> { + let app = App::new() + .register(config::plugin()) + .register(your_routes::plugin()) // Add your plugin here + .init() + .await?; + + Ok(app) +} +``` + +## Middleware Plugins + +The template includes a `security` plugin by default. It adds common response headers, as well as a request body limiter and timeout using `tower::ServiceBuilder` and `tower-http`. + +When request logging is enabled during generation, the template also includes a `logging` plugin that adds request IDs and request/response tracing. + +## Error Handling + +Route handlers can return `AppResult`, which is an alias for `Result`. `AppError` implements `IntoResponse`, so API failures are returned as JSON. It also implements `From`, so handlers can use `?` with `anyhow` errors: + +```rust +use anyhow::Context; + +use crate::error::AppResult; + +async fn handler() -> AppResult { + do_work().await.context("failed to do work")?; + Ok("done".to_string()) +} +``` + +## License + +Configure your license as needed. diff --git a/server-new/src/config.rs b/server-new/src/config.rs new file mode 100644 index 0000000..e88390b --- /dev/null +++ b/server-new/src/config.rs @@ -0,0 +1,78 @@ +use std::net::{IpAddr, Ipv4Addr}; + +use anyhow::Context; +use axum_plugin::AdHocPlugin; +use serde::Deserialize; + +use crate::state::AppState; + +/// Parsed app configuration +#[derive(Debug, Clone, Deserialize)] +pub struct AppConfig { + // Server config + #[serde(default = "default_host")] + pub host: IpAddr, + #[serde(default = "default_port")] + pub port: u16, + #[serde(default = "default_log_level")] + pub log_level: String, + #[serde(default = "default_request_id_header")] + pub request_id_header: String, + + // Auth + pub cookie_key: String, + #[serde(default = "default_cookie_name")] + pub cookie_name: String, + #[serde(default = "default_session_length")] + pub session_length: i64, + + // Security + #[serde(default = "default_body_limit")] + pub body_limit: usize, + #[serde(default = "default_req_timeout")] + pub request_timeout: u64, +} +fn default_host() -> IpAddr { + IpAddr::V4(Ipv4Addr::LOCALHOST) +} +fn default_port() -> u16 { + 8080 +} +fn default_log_level() -> String { + "info".to_string() +} +fn default_request_id_header() -> String { + "x-request-id".to_string() +} +fn default_cookie_name() -> String { + "auth-rs-chat".to_string() +} +fn default_session_length() -> i64 { + 60 * 60 * 24 * 7 // 1 week +} + +fn default_body_limit() -> usize { + 2 * 1024 * 1024 // 2 MB +} +fn default_req_timeout() -> u64 { + 120 // 2 minutes +} + +/// Plugin that reads and validates configuration, and adds it to server state +pub fn plugin() -> AdHocPlugin { + AdHocPlugin::named("Config").on_init(async |mut state| { + let config = extract_config()?; + state.insert(config); + Ok(state) + }) +} + +/// Extract the configuration from env variables prefixed with `RS_CHAT_`. +fn extract_config() -> anyhow::Result { + let config = figment::Figment::new() + .merge(figment::providers::Env::prefixed("RS_CHAT_")) + .extract::() + .context("Failed to extract valid configuration")?; + + Ok(config) +} diff --git a/server-new/src/error.rs b/server-new/src/error.rs new file mode 100644 index 0000000..c0b2e0f --- /dev/null +++ b/server-new/src/error.rs @@ -0,0 +1,88 @@ +use axum::{ + Json, + http::StatusCode, + response::{IntoResponse, Response}, +}; +use serde::Serialize; + +/// Global result type that can be used for API route handlers +pub type AppResult = Result; + +/// Global error type +#[derive(Debug)] +pub struct AppError { + status: StatusCode, + message: String, + source: Option, +} + +impl From for AppError { + fn from(error: anyhow::Error) -> Self { + Self::internal(error) + } +} + +// Add more conversions here to be able to propagate them in route handlers: +// impl From for AppError { +// fn from(error: DatabaseError) -> Self { +// Self::internal(error.into()) +// } +// } + +impl AppError { + pub fn new(status: StatusCode, message: impl Into) -> Self { + Self { + status, + message: message.into(), + source: None, + } + } + + pub fn bad_request(message: impl Into) -> Self { + Self::new(StatusCode::BAD_REQUEST, message) + } + + pub fn unauthorized() -> Self { + Self::new(StatusCode::UNAUTHORIZED, "unauthorized") + } + + pub fn not_found(message: impl Into) -> Self { + Self::new(StatusCode::NOT_FOUND, message) + } + + pub fn internal(error: anyhow::Error) -> Self { + Self { + status: StatusCode::INTERNAL_SERVER_ERROR, + message: "internal server error".to_string(), + source: Some(error), + } + } +} + +#[derive(Debug, Serialize)] +struct ErrorResponse { + error: ErrorBody, +} + +#[derive(Debug, Serialize)] +struct ErrorBody { + message: String, + status: u16, +} + +impl IntoResponse for AppError { + fn into_response(self) -> Response { + if let Some(error) = self.source { + tracing::warn!(error = ?error, "request failed"); + } + + let response = ErrorResponse { + error: ErrorBody { + message: self.message, + status: self.status.as_u16(), + }, + }; + + (self.status, Json(response)).into_response() + } +} diff --git a/server-new/src/extractors/mod.rs b/server-new/src/extractors/mod.rs new file mode 100644 index 0000000..f52f1c4 --- /dev/null +++ b/server-new/src/extractors/mod.rs @@ -0,0 +1 @@ +pub mod session; diff --git a/server-new/src/extractors/session.rs b/server-new/src/extractors/session.rs new file mode 100644 index 0000000..4ced4aa --- /dev/null +++ b/server-new/src/extractors/session.rs @@ -0,0 +1,70 @@ +use anyhow::Context; +use axum::extract::{FromRequestParts, OptionalFromRequestParts}; +use serde::{Deserialize, Serialize}; +use tower_sessions::Session; + +use crate::error::AppError; + +/// The field in the `tower_sessions` storage used to store the user session data +const USER_SESSION_FIELD: &str = "sess"; + +/// Active session data. Beware when changing or adding to this struct, as it can +/// invalidate existing sessions. +/// +/// This can be used as an extractor in route handlers: +/// - If used as `Option`, will be `Some` if there is an active session +/// and `None` otherwise. +/// - If used as `UserSession`, request will automatically return an unauthorized error +/// if there is no active session. +#[derive(Debug, Serialize, Deserialize)] +pub struct UserSession { + pub user_id: String, +} + +impl UserSession { + pub fn new(user_id: String) -> Self { + Self { user_id } + } + + pub async fn init(session: &Session, user_id: &str) -> Result<(), AppError> { + Ok(session + .insert(USER_SESSION_FIELD, UserSession::new(user_id.into())) + .await + .context("failed to initialize session")?) + } +} + +impl OptionalFromRequestParts for UserSession { + type Rejection = AppError; + + async fn from_request_parts( + parts: &mut axum::http::request::Parts, + state: &S, + ) -> Result, Self::Rejection> { + let session = Session::from_request_parts(parts, state) + .await + .map_err(|(_, msg)| AppError::internal(anyhow::anyhow!(msg)))?; + + match session.get::(USER_SESSION_FIELD).await { + Ok(Some(user_session)) => Ok(Some(user_session)), + Ok(None) => Ok(None), + Err(err) => Err(AppError::internal( + anyhow::Error::from(err).context("error while retrieving session"), + )), + } + } +} + +impl FromRequestParts for UserSession { + type Rejection = AppError; + + async fn from_request_parts( + parts: &mut axum::http::request::Parts, + state: &S, + ) -> Result { + match >::from_request_parts(parts, state).await? { + Some(user_session) => Ok(user_session), + None => Err(AppError::unauthorized()), + } + } +} diff --git a/server-new/src/lib.rs b/server-new/src/lib.rs new file mode 100644 index 0000000..7de7e9b --- /dev/null +++ b/server-new/src/lib.rs @@ -0,0 +1,23 @@ +use axum_plugin::{App, InitializedApp}; + +use crate::state::AppState; + +mod config; +mod error; +mod extractors; +mod plugins; +mod routes; +mod state; + +pub async fn create_app() -> anyhow::Result> { + let app = App::new() + .register(config::plugin()) // Extract configuration and add to state + .register(routes::plugin()) // Add API routes + .register(plugins::session::plugin()) // Setup sessions + .register(plugins::logging::plugin()) // Request logging + .register(plugins::security::plugin()) // Body limit, security headers, etc. + .init() + .await?; + + Ok(app) +} diff --git a/server-new/src/main.rs b/server-new/src/main.rs new file mode 100644 index 0000000..2907e3e --- /dev/null +++ b/server-new/src/main.rs @@ -0,0 +1,96 @@ +use std::net::SocketAddr; + +use rs_chat_api::create_app; +use tracing::level_filters::LevelFilter; +use tracing_subscriber::{EnvFilter, Registry, layer::SubscriberExt, util::SubscriberInitExt}; + +#[tokio::main] +async fn main() -> anyhow::Result<()> { + // Read .env file in debug mode + #[cfg(debug_assertions)] + dotenvy::dotenv().ok(); + + // Initialize logging + let (log_filter_handle, _log_guard) = init_logging(); + + // Build server + let app = create_app().await?; + let config = &app.state().config; + + // Set log level from config + let env_filter = tracing_subscriber::EnvFilter::builder() + .with_default_directive(LevelFilter::INFO.into()) + .parse(&config.log_level)?; + log_filter_handle.reload(env_filter)?; + + // Start listening for requests + let addr = SocketAddr::new(config.host, config.port); + let listener = tokio::net::TcpListener::bind(addr).await?; + tracing::info!("Server listening on http://{}...", listener.local_addr()?); + axum::serve(listener, app.router().into_make_service()) + .with_graceful_shutdown(shutdown_signal(app.shutdown())) + .await?; + + Ok(()) +} + +fn init_logging() -> ( + tracing_subscriber::reload::Handle, + tracing_appender::non_blocking::WorkerGuard, +) { + let init_log_level = if cfg!(debug_assertions) { + dotenvy::var("RS_CHAT_LOG_LEVEL").unwrap_or("info".into()) + } else { + "info".to_owned() + }; + let (writer, guard) = tracing_appender::non_blocking(std::io::stdout()); + let (filter_layer, filter_handle) = + tracing_subscriber::reload::Layer::new(EnvFilter::new(init_log_level)); + + if cfg!(debug_assertions) { + tracing_subscriber::registry() + .with(filter_layer) + .with(tracing_subscriber::fmt::layer().with_writer(writer)) + .init(); + } else { + let json_layer = tracing_subscriber::fmt::layer() + .json() + .flatten_event(true) + .with_current_span(false) + .with_writer(writer); + + tracing_subscriber::registry() + .with(filter_layer) + .with(json_layer) + .init(); + } + + (filter_handle, guard) +} + +/// Shutdown signal: listens for Ctrl-C, SIGINT, SIGTERM signals +async fn shutdown_signal(on_shutdown: impl Future + Send) { + let ctrl_c = async { + tokio::signal::ctrl_c() + .await + .expect("failed to register Ctrl-C handler"); + }; + + #[cfg(unix)] + let terminate = async { + tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) + .expect("failed to register SIGTERM handler") + .recv() + .await; + }; + + #[cfg(not(unix))] + let terminate = std::future::pending::<()>(); + + tokio::select! { + _ = ctrl_c => {}, + _ = terminate => {}, + } + tracing::info!("Received shutdown signal, shutting down server..."); + on_shutdown.await; +} diff --git a/server-new/src/plugins/logging.rs b/server-new/src/plugins/logging.rs new file mode 100644 index 0000000..bdd85d2 --- /dev/null +++ b/server-new/src/plugins/logging.rs @@ -0,0 +1,49 @@ +use std::str::FromStr; + +use anyhow::Context; +use axum::{extract::Request, http::HeaderName}; +use axum_plugin::AdHocPlugin; +use tower::ServiceBuilder; +use tower_http::{ + request_id::{MakeRequestUuid, PropagateRequestIdLayer, SetRequestIdLayer}, + trace::{DefaultOnRequest, DefaultOnResponse, TraceLayer}, +}; +use tracing::{Level, level_filters::LevelFilter}; + +use crate::state::AppState; + +pub fn plugin() -> AdHocPlugin { + AdHocPlugin::named("Request logs").on_setup(|router, state: &AppState| { + if LevelFilter::from_str(&state.config.log_level)? > LevelFilter::INFO { + return Ok(router); + } + + const LOG_LEVEL: Level = Level::INFO; + let request_id_header = HeaderName::from_str(&state.config.request_id_header) + .context("invalid request ID header")?; + + let trace_layer = TraceLayer::new_for_http() + .make_span_with({ + let id_header = request_id_header.clone(); + move |req: &Request| { + tracing::span!(LOG_LEVEL, "request", + method = %req.method(), + uri = %req.uri(), + id = req.headers().get(&id_header).and_then(|id| id.to_str().ok()), + ) + } + }) + .on_request(DefaultOnRequest::new().level(LOG_LEVEL)) + .on_response(DefaultOnResponse::new().level(LOG_LEVEL)); + + let logging_service = ServiceBuilder::new() + .layer(SetRequestIdLayer::new( + request_id_header.clone(), + MakeRequestUuid::default(), + )) + .layer(trace_layer) + .layer(PropagateRequestIdLayer::new(request_id_header)); + + Ok(router.layer(logging_service)) + }) +} diff --git a/server-new/src/plugins/mod.rs b/server-new/src/plugins/mod.rs new file mode 100644 index 0000000..e40cded --- /dev/null +++ b/server-new/src/plugins/mod.rs @@ -0,0 +1,3 @@ +pub mod logging; +pub mod security; +pub mod session; diff --git a/server-new/src/plugins/security.rs b/server-new/src/plugins/security.rs new file mode 100644 index 0000000..9b7f0d9 --- /dev/null +++ b/server-new/src/plugins/security.rs @@ -0,0 +1,32 @@ +use std::time::Duration; + +use axum::http::StatusCode; +use axum_plugin::AdHocPlugin; +use tower::ServiceBuilder; +use tower_http::{limit::RequestBodyLimitLayer, timeout::TimeoutLayer}; + +use crate::state::AppState; + +/// # Security plugin +/// Includes body limiter, request timeout, and security headers. +pub fn plugin() -> AdHocPlugin { + AdHocPlugin::named("Security").on_setup(|router, state: &AppState| { + let security_headers = axum_helmet::Helmet::new() + .add(axum_helmet::CrossOriginOpenerPolicy::same_origin()) + .add(axum_helmet::CrossOriginResourcePolicy::same_origin()) + .add(axum_helmet::ReferrerPolicy::no_referrer()) + .add(axum_helmet::XContentTypeOptions::nosniff()) + .add(axum_helmet::XFrameOptions::same_origin()) + .into_layer()?; + + let service = ServiceBuilder::new() + .layer(RequestBodyLimitLayer::new(state.config.body_limit)) + .layer(TimeoutLayer::with_status_code( + StatusCode::REQUEST_TIMEOUT, + Duration::from_secs(state.config.request_timeout), + )) + .layer(security_headers); + + Ok(router.layer(service)) + }) +} diff --git a/server-new/src/plugins/session.rs b/server-new/src/plugins/session.rs new file mode 100644 index 0000000..1fff624 --- /dev/null +++ b/server-new/src/plugins/session.rs @@ -0,0 +1,33 @@ +use anyhow::{Context, bail}; +use axum_plugin::AdHocPlugin; +use tower_sessions::{ + Expiry, SessionManagerLayer, + cookie::{Key, SameSite, time::Duration}, +}; + +use crate::state::AppState; + +pub fn plugin() -> AdHocPlugin { + AdHocPlugin::named("Session").on_setup(|router, state: &AppState| { + let cookie_key = + hex::decode(&state.config.cookie_key).context("cookie_key must be hex value")?; + if cookie_key.len() < 32 { + bail!("cookie_key must be at least 32 bytes"); + } + + // TODO change to a persistent store! + let session_store = tower_sessions::MemoryStore::default(); + let session_layer = SessionManagerLayer::new(session_store) + .with_name(state.config.cookie_name.clone()) + .with_expiry(Expiry::OnInactivity(Duration::seconds( + state.config.session_length, + ))) + .with_private(Key::derive_from(&cookie_key)) + .with_path("/") + .with_secure(true) + .with_http_only(true) + .with_same_site(SameSite::Lax); + + Ok(router.layer(session_layer)) + }) +} diff --git a/server-new/src/routes/auth.rs b/server-new/src/routes/auth.rs new file mode 100644 index 0000000..defeef4 --- /dev/null +++ b/server-new/src/routes/auth.rs @@ -0,0 +1,36 @@ +use anyhow::Context; +use axum::response::IntoResponse; + +use crate::{error::AppResult, extractors::session::UserSession, state::AppState}; + +pub fn routes() -> axum::Router { + axum::Router::new() + .route("/login", axum::routing::post(login_handler)) + .route("/user", axum::routing::get(get_user_handler)) + .route("/logout", axum::routing::post(logout_handler)) +} + +async fn login_handler(session: tower_sessions::Session) -> AppResult { + // TODO login handling logic + let user_id = "user123"; + UserSession::init(&session, user_id).await?; + + Ok(format!("Logged in as {user_id}")) +} + +async fn get_user_handler(UserSession { user_id }: UserSession) -> impl IntoResponse { + format!("Logged in as {user_id}") +} + +async fn logout_handler( + maybe_user: Option, + session: tower_sessions::Session, +) -> AppResult { + match maybe_user { + Some(_) => { + session.delete().await.context("error logging out")?; + Ok("Logged out") + } + None => Ok("Already logged out"), + } +} diff --git a/server-new/src/routes/health.rs b/server-new/src/routes/health.rs new file mode 100644 index 0000000..5998a39 --- /dev/null +++ b/server-new/src/routes/health.rs @@ -0,0 +1,9 @@ +use crate::state::AppState; + +pub fn routes() -> axum::Router { + axum::Router::new().route("/", axum::routing::get(health_handler)) +} + +async fn health_handler() -> &'static str { + "OK" +} diff --git a/server-new/src/routes/hello.rs b/server-new/src/routes/hello.rs new file mode 100644 index 0000000..52fd573 --- /dev/null +++ b/server-new/src/routes/hello.rs @@ -0,0 +1,15 @@ +use crate::{error::AppResult, state::AppState}; + +pub fn routes() -> axum::Router { + axum::Router::new() + .route("/", axum::routing::get(hello_handler)) + .route("/", axum::routing::post(post_handler)) +} + +async fn hello_handler() -> AppResult { + Ok("Hello, World!".to_string()) +} + +async fn post_handler() -> AppResult { + Ok("Post handler!".to_string()) +} diff --git a/server-new/src/routes/mod.rs b/server-new/src/routes/mod.rs new file mode 100644 index 0000000..f1662ca --- /dev/null +++ b/server-new/src/routes/mod.rs @@ -0,0 +1,19 @@ +use axum_plugin::AdHocPlugin; + +use crate::state::AppState; + +pub mod auth; +pub mod health; +pub mod hello; + +/// Adds all API routes to the server under `/api` +pub fn plugin() -> AdHocPlugin { + AdHocPlugin::named("API routes").on_setup(|router, _state| { + let api_routes = axum::Router::new() + .nest("/auth", auth::routes()) + .nest("/hello", hello::routes()) + .nest("/health", health::routes()); + + Ok(router.nest("/api", api_routes)) + }) +} diff --git a/server-new/src/state.rs b/server-new/src/state.rs new file mode 100644 index 0000000..6b9a0e8 --- /dev/null +++ b/server-new/src/state.rs @@ -0,0 +1,33 @@ +//! Application state + +use std::{ops::Deref, sync::Arc}; + +use axum_plugin::{AppState, TypeMap}; + +use crate::config::AppConfig; + +/// App state stored in the Axum router +#[derive(Clone)] +pub struct AppState(Arc); + +#[derive(AppState)] +pub struct AppStateInner { + pub config: AppConfig, + // add state here... +} + +impl Deref for AppState { + type Target = AppStateInner; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl TryFrom for AppState { + type Error = anyhow::Error; + + fn try_from(map: TypeMap) -> Result { + Ok(Self(Arc::new(AppStateInner::try_from(map)?))) + } +} From 427beb4a5ab475eed8171b2d24f40afac2ccf72d Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 22 Jun 2026 16:36:05 -0400 Subject: [PATCH 022/111] add database --- server-new/.env.example | 3 + server-new/Cargo.lock | 526 +++++++++++++++++- server-new/Cargo.toml | 10 +- .../down.sql | 6 + .../up.sql | 36 ++ .../down.sql | 5 + .../up.sql | 25 + .../2025-06-15-031207_add_users/down.sql | 6 + .../2025-06-15-031207_add_users/up.sql | 15 + .../2025-06-16-035815_add_api_keys/down.sql | 3 + .../2025-06-16-035815_add_api_keys/up.sql | 12 + .../down.sql | 2 + .../2025-06-17-063104_add_message_meta/up.sql | 2 + .../down.sql | 10 + .../2025-06-21-031403_full_text_search/up.sql | 53 ++ .../down.sql | 3 + .../up.sql | 11 + .../down.sql | 16 + .../up.sql | 16 + .../down.sql | 1 + .../2025-07-11-012329_add_app_api_keys/up.sql | 8 + .../2025-07-13-170127_add_tools/down.sql | 13 + .../2025-07-13-170127_add_tools/up.sql | 26 + .../2025-07-17-223807_add_providers/down.sql | 7 + .../2025-07-17-223807_add_providers/up.sql | 68 +++ .../down.sql | 2 + .../2025-08-08-080101_add_session_meta/up.sql | 2 + .../down.sql | 4 + .../up.sql | 4 + .../2025-08-16-152113_refactor_tools/down.sql | 3 + .../2025-08-16-152113_refactor_tools/up.sql | 23 + .../2025-08-31-034235_add_files/down.sql | 1 + .../2025-08-31-034235_add_files/up.sql | 18 + .../down.sql | 15 + .../2025-09-03-063406_remove_old_tools/up.sql | 2 + server-new/src/config.rs | 9 +- server-new/src/db/mod.rs | 34 ++ server-new/src/db/models.rs | 17 + server-new/src/db/models/user.rs | 49 ++ server-new/src/db/repositories/mod.rs | 3 + server-new/src/db/repositories/user.rs | 117 ++++ server-new/src/db/schema.rs | 146 +++++ server-new/src/error.rs | 11 +- server-new/src/extractors/session.rs | 80 ++- server-new/src/lib.rs | 3 + server-new/src/main.rs | 16 +- server-new/src/plugins/database.rs | 51 ++ server-new/src/plugins/logging.rs | 6 +- server-new/src/plugins/mod.rs | 1 + server-new/src/routes/auth.rs | 31 +- server-new/src/services/mod.rs | 3 + server-new/src/services/user.rs | 23 + server-new/src/state.rs | 10 +- 53 files changed, 1522 insertions(+), 45 deletions(-) create mode 100644 server-new/migrations/00000000000000_diesel_initial_setup/down.sql create mode 100644 server-new/migrations/00000000000000_diesel_initial_setup/up.sql create mode 100644 server-new/migrations/2025-06-12-171524_create_chat_sessions/down.sql create mode 100644 server-new/migrations/2025-06-12-171524_create_chat_sessions/up.sql create mode 100644 server-new/migrations/2025-06-15-031207_add_users/down.sql create mode 100644 server-new/migrations/2025-06-15-031207_add_users/up.sql create mode 100644 server-new/migrations/2025-06-16-035815_add_api_keys/down.sql create mode 100644 server-new/migrations/2025-06-16-035815_add_api_keys/up.sql create mode 100644 server-new/migrations/2025-06-17-063104_add_message_meta/down.sql create mode 100644 server-new/migrations/2025-06-17-063104_add_message_meta/up.sql create mode 100644 server-new/migrations/2025-06-21-031403_full_text_search/down.sql create mode 100644 server-new/migrations/2025-06-21-031403_full_text_search/up.sql create mode 100644 server-new/migrations/2025-06-23-023453_update_session_updated_at/down.sql create mode 100644 server-new/migrations/2025-06-23-023453_update_session_updated_at/up.sql create mode 100644 server-new/migrations/2025-07-10-165227_add_auth_providers/down.sql create mode 100644 server-new/migrations/2025-07-10-165227_add_auth_providers/up.sql create mode 100644 server-new/migrations/2025-07-11-012329_add_app_api_keys/down.sql create mode 100644 server-new/migrations/2025-07-11-012329_add_app_api_keys/up.sql create mode 100644 server-new/migrations/2025-07-13-170127_add_tools/down.sql create mode 100644 server-new/migrations/2025-07-13-170127_add_tools/up.sql create mode 100644 server-new/migrations/2025-07-17-223807_add_providers/down.sql create mode 100644 server-new/migrations/2025-07-17-223807_add_providers/up.sql create mode 100644 server-new/migrations/2025-08-08-080101_add_session_meta/down.sql create mode 100644 server-new/migrations/2025-08-08-080101_add_session_meta/up.sql create mode 100644 server-new/migrations/2025-08-09-160537_remove_unused_secret_field/down.sql create mode 100644 server-new/migrations/2025-08-09-160537_remove_unused_secret_field/up.sql create mode 100644 server-new/migrations/2025-08-16-152113_refactor_tools/down.sql create mode 100644 server-new/migrations/2025-08-16-152113_refactor_tools/up.sql create mode 100644 server-new/migrations/2025-08-31-034235_add_files/down.sql create mode 100644 server-new/migrations/2025-08-31-034235_add_files/up.sql create mode 100644 server-new/migrations/2025-09-03-063406_remove_old_tools/down.sql create mode 100644 server-new/migrations/2025-09-03-063406_remove_old_tools/up.sql create mode 100644 server-new/src/db/mod.rs create mode 100644 server-new/src/db/models.rs create mode 100644 server-new/src/db/models/user.rs create mode 100644 server-new/src/db/repositories/mod.rs create mode 100644 server-new/src/db/repositories/user.rs create mode 100644 server-new/src/db/schema.rs create mode 100644 server-new/src/plugins/database.rs create mode 100644 server-new/src/services/mod.rs create mode 100644 server-new/src/services/user.rs diff --git a/server-new/.env.example b/server-new/.env.example index c83e520..ade296f 100644 --- a/server-new/.env.example +++ b/server-new/.env.example @@ -2,6 +2,9 @@ RS_CHAT_HOST=127.0.0.1 RS_CHAT_PORT=8080 +# Database +RS_CHAT_DATABASE_URL=postgres://postgres:postgres@localhost/postgres + # Auth RS_CHAT_COOKIE_KEY= # hex secret >=32 bytes, e.g. `openssl rand --hex 32` diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index c82280c..b909299 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -198,6 +198,12 @@ version = "1.25.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c8efb64bd706a16a1bdde310ae86b351e4d21550d98d056f22f8a7f7a2183fec" +[[package]] +name = "byteorder" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b" + [[package]] name = "bytes" version = "1.12.0" @@ -282,6 +288,58 @@ dependencies = [ "cipher", ] +[[package]] +name = "darling" +version = "0.21.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9cdf337090841a411e2a7f3deb9187445851f91b309c0c0a29e05f74a00a48c0" +dependencies = [ + "darling_core", + "darling_macro", +] + +[[package]] +name = "darling_core" +version = "0.21.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1247195ecd7e3c85f83c8d2a366e4210d588e802133e1e355180a9870b517ea4" +dependencies = [ + "fnv", + "ident_case", + "proc-macro2", + "quote", + "strsim", + "syn", +] + +[[package]] +name = "darling_macro" +version = "0.21.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d38308df82d1080de0afee5d069fa14b0326a88c14f15c5ccda35b4a6c414c81" +dependencies = [ + "darling_core", + "quote", + "syn", +] + +[[package]] +name = "deadpool" +version = "0.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "883466cb8db62725aee5f4a6011e8a5d42912b42632df32aad57fc91127c6e04" +dependencies = [ + "deadpool-runtime", + "num_cpus", + "tokio", +] + +[[package]] +name = "deadpool-runtime" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2657f61fb1dd8bf37a8d51093cc7cee4e77125b22f7753f49b289f831bec2bae" + [[package]] name = "deranged" version = "0.5.8" @@ -291,6 +349,70 @@ dependencies = [ "serde_core", ] +[[package]] +name = "diesel" +version = "2.3.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29fe29a87fb84c631ffb3ba21798c4b1f3a964701ba78f0dce4bf8668562ec88" +dependencies = [ + "bitflags", + "byteorder", + "diesel_derives", + "downcast-rs", + "itoa", + "pq-sys", + "time", + "uuid", +] + +[[package]] +name = "diesel-async" +version = "0.9.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dd39af30158d444884f166fe4c58f35dc40ad71ad017bb59408a3448526ff4bd" +dependencies = [ + "deadpool", + "diesel", + "futures-core", + "futures-util", + "pin-project-lite", + "tokio", + "tokio-postgres", +] + +[[package]] +name = "diesel_derives" +version = "2.3.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d1817b7f4279b947fc4cafddec12b0e5f8727141706561ce3ac94a60bddd1cf5" +dependencies = [ + "diesel_table_macro_syntax", + "dsl_auto_type", + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "diesel_migrations" +version = "2.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "28d0f4a98124ba6d4ca75da535f65984badec16a003b6e2f94a01e31a79490b8" +dependencies = [ + "diesel", + "migrations_internals", + "migrations_macros", +] + +[[package]] +name = "diesel_table_macro_syntax" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fe2444076b48641147115697648dc743c2c00b61adade0f01ce67133c7babe8c" +dependencies = [ + "syn", +] + [[package]] name = "digest" version = "0.10.7" @@ -308,6 +430,32 @@ version = "0.15.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1aaf95b3e5c8f23aa320147307562d361db0ae0d51242340f558153b4eb2439b" +[[package]] +name = "downcast-rs" +version = "2.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "117240f60069e65410b3ae1bb213295bd828f707b5bec6596a1afc8793ce0cbc" + +[[package]] +name = "dsl_auto_type" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dd122633e4bef06db27737f21d3738fb89c8f6d5360d6d9d7635dda142a7757e" +dependencies = [ + "darling", + "either", + "heck", + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "either" +version = "1.16.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" + [[package]] name = "errno" version = "0.3.14" @@ -318,6 +466,12 @@ dependencies = [ "windows-sys", ] +[[package]] +name = "fallible-iterator" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4443176a9f2c162692bd3d352d745ef9413eec5782a80d8fd6f8a1ac692a07f7" + [[package]] name = "figment" version = "0.10.19" @@ -331,6 +485,12 @@ dependencies = [ "version_check", ] +[[package]] +name = "fnv" +version = "1.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" + [[package]] name = "form_urlencoded" version = "1.2.2" @@ -446,7 +606,7 @@ checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0" dependencies = [ "cfg-if", "libc", - "wasi", + "wasi 0.11.1+wasi-snapshot-preview1", ] [[package]] @@ -482,12 +642,24 @@ dependencies = [ "polyval", ] +[[package]] +name = "heck" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" + [[package]] name = "helmet-core" version = "1.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8a865b9c8b67316ab132710af828252e764cdf2195bbcd72a23b96127150d9de" +[[package]] +name = "hermit-abi" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" + [[package]] name = "hex" version = "0.4.3" @@ -592,6 +764,12 @@ dependencies = [ "tower-service", ] +[[package]] +name = "ident_case" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9e0384b61958566e926dc50660321d12159025e767c18e043daf26b70104c39" + [[package]] name = "inlinable_string" version = "0.1.15" @@ -636,6 +814,15 @@ version = "0.2.186" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" +[[package]] +name = "libredox" +version = "0.1.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e02f3bb43d335493c96bf3fd3a321600bf6bd07ed34bc64118e9293bdffea46c" +dependencies = [ + "libc", +] + [[package]] name = "lock_api" version = "0.4.14" @@ -666,12 +853,43 @@ version = "0.8.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" +[[package]] +name = "md-5" +version = "0.10.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d89e7ee0cfbedfc4da3340218492196241d89eefb6dab27de5df917a6d2e78cf" +dependencies = [ + "cfg-if", + "digest", +] + [[package]] name = "memchr" version = "2.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "88904434abc2901f197fe8cc55f0445e7ded921dba5911dad2e2b39b48e663c4" +[[package]] +name = "migrations_internals" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "36c791ecdf977c99f45f23280405d7723727470f6689a5e6dbf513ac547ae10d" +dependencies = [ + "serde", + "toml", +] + +[[package]] +name = "migrations_macros" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "36fc5ac76be324cfd2d3f2cf0fdf5d5d3c4f14ed8aaebadb09e304ba42282703" +dependencies = [ + "migrations_internals", + "proc-macro2", + "quote", +] + [[package]] name = "mime" version = "0.3.17" @@ -685,7 +903,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "02bd0af71c67b473010cbbc60715ee815645a4dc942899111f494b4b737d6fda" dependencies = [ "libc", - "wasi", + "wasi 0.11.1+wasi-snapshot-preview1", "windows-sys", ] @@ -704,6 +922,34 @@ version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "521739c6d2bac4aa25192232afe6841231376b2b26d4d9fae5ecf8ca5772e441" +[[package]] +name = "num_cpus" +version = "1.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91df4bbde75afed763b708b7eee1e8e7651e02d97f6d5dd763e89367e957b23b" +dependencies = [ + "hermit-abi", + "libc", +] + +[[package]] +name = "objc2-core-foundation" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2a180dd8642fa45cdb7dd721cd4c11b1cadd4929ce112ebd8b9f5803cc79d536" +dependencies = [ + "bitflags", +] + +[[package]] +name = "objc2-system-configuration" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7216bd11cbda54ccabcab84d523dc93b858ec75ecfb3a7d89513fa22464da396" +dependencies = [ + "objc2-core-foundation", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -768,12 +1014,37 @@ version = "2.3.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" +[[package]] +name = "phf" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c1562dc717473dbaa4c1f85a36410e03c047b2e7df7f45ee938fbef64ae7fadf" +dependencies = [ + "phf_shared", + "serde", +] + +[[package]] +name = "phf_shared" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e57fef6bc5981e38c2ce2d63bfa546861309f875b8a75f092d1d54ae2d64f266" +dependencies = [ + "siphasher", +] + [[package]] name = "pin-project-lite" version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" +[[package]] +name = "pkg-config" +version = "0.3.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" + [[package]] name = "polyval" version = "0.6.2" @@ -786,6 +1057,35 @@ dependencies = [ "universal-hash", ] +[[package]] +name = "postgres-protocol" +version = "0.6.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3ee9dd5fe15055d2b6806f4736aa0c9637217074e224bbec46d4041b91bb9491" +dependencies = [ + "base64", + "byteorder", + "bytes", + "fallible-iterator", + "hmac", + "md-5", + "memchr", + "rand 0.9.4", + "sha2", + "stringprep", +] + +[[package]] +name = "postgres-types" +version = "0.2.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "54b858f82211e84682fecd373f68e1ceae642d8d751a1ebd13f33de6257b3e20" +dependencies = [ + "bytes", + "fallible-iterator", + "postgres-protocol", +] + [[package]] name = "powerfmt" version = "0.2.0" @@ -801,6 +1101,17 @@ dependencies = [ "zerocopy", ] +[[package]] +name = "pq-sys" +version = "0.7.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "574ddd6a267294433f140b02a726b0640c43cf7c6f717084684aaa3b285aba61" +dependencies = [ + "libc", + "pkg-config", + "vcpkg", +] + [[package]] name = "proc-macro2" version = "1.0.106" @@ -937,6 +1248,9 @@ dependencies = [ "axum", "axum-helmet", "axum-plugin", + "diesel", + "diesel-async", + "diesel_migrations", "dotenvy", "figment", "hex", @@ -950,6 +1264,7 @@ dependencies = [ "tracing", "tracing-appender", "tracing-subscriber", + "uuid", ] [[package]] @@ -1030,6 +1345,15 @@ dependencies = [ "serde_core", ] +[[package]] +name = "serde_spanned" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6662b5879511e06e8999a8a235d848113e942c9124f211511b16466ee2995f26" +dependencies = [ + "serde_core", +] + [[package]] name = "serde_urlencoded" version = "0.7.1" @@ -1072,6 +1396,12 @@ dependencies = [ "libc", ] +[[package]] +name = "siphasher" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b2aa850e253778c88a04c3d7323b043aeda9d3e30d5971937c1855769763678e" + [[package]] name = "slab" version = "0.4.12" @@ -1094,6 +1424,23 @@ dependencies = [ "windows-sys", ] +[[package]] +name = "stringprep" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b4df3d392d81bd458a8a621b8bffbd2302a12ffe288a9d931670948749463b1" +dependencies = [ + "unicode-bidi", + "unicode-normalization", + "unicode-properties", +] + +[[package]] +name = "strsim" +version = "0.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" + [[package]] name = "subtle" version = "2.6.1" @@ -1154,9 +1501,9 @@ dependencies = [ [[package]] name = "time" -version = "0.3.49" +version = "0.3.51" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "711a53c2d47bbd818258c498c8dbfe186a2526c631495cfe7e078567f86b8469" +checksum = "85c17d80feb7334b40c484e45ed1a5273dfd8bfda537c3be2e74a06a6686f327" dependencies = [ "deranged", "num-conv", @@ -1174,20 +1521,36 @@ checksum = "9e1c906769ad99c88eaa54e728060edef082f8e358ff32030cb7c7d315e81109" [[package]] name = "time-macros" -version = "0.2.29" +version = "0.2.30" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "71c652a3727a9cbb9a02f707f530b618ce00d0ccd762009c8c23bd191df3c17d" +checksum = "dcef1a61bdb119096e153208ec5cbec23944ce8bca13be5c7f60c634f7403935" dependencies = [ "num-conv", "time-core", ] +[[package]] +name = "tinyvec" +version = "1.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3e61e67053d25a4e82c844e8424039d9745781b3fc4f32b8d55ed50f5f667ef3" +dependencies = [ + "tinyvec_macros", +] + +[[package]] +name = "tinyvec_macros" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" + [[package]] name = "tokio" version = "1.52.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe" dependencies = [ + "bytes", "libc", "mio", "pin-project-lite", @@ -1208,6 +1571,76 @@ dependencies = [ "syn", ] +[[package]] +name = "tokio-postgres" +version = "0.7.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dcea47c8f71744367793f16c2db1f11cb859d28f436bdb4ca9193eb1f787ee42" +dependencies = [ + "async-trait", + "byteorder", + "bytes", + "fallible-iterator", + "futures-channel", + "futures-util", + "log", + "parking_lot", + "percent-encoding", + "phf", + "pin-project-lite", + "postgres-protocol", + "postgres-types", + "rand 0.9.4", + "socket2", + "tokio", + "tokio-util", + "whoami", +] + +[[package]] +name = "tokio-util" +version = "0.7.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ae9cec805b01e8fc3fd2fe289f89149a9b66dd16786abd8b19cfa7b48cb0098" +dependencies = [ + "bytes", + "futures-core", + "futures-sink", + "pin-project-lite", + "tokio", +] + +[[package]] +name = "toml" +version = "0.9.12+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf92845e79fc2e2def6a5d828f0801e29a2f8acc037becc5ab08595c7d5e9863" +dependencies = [ + "serde_core", + "serde_spanned", + "toml_datetime", + "toml_parser", + "winnow 0.7.15", +] + +[[package]] +name = "toml_datetime" +version = "0.7.5+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92e1cfed4a3038bc5a127e35a2d360f145e1f4b971b551a2ba5fd7aedf7e1347" +dependencies = [ + "serde_core", +] + +[[package]] +name = "toml_parser" +version = "1.1.2+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a2abe9b86193656635d2411dc43050282ca48aa31c2451210f4202550afb7526" +dependencies = [ + "winnow 1.0.3", +] + [[package]] name = "tower" version = "0.5.3" @@ -1435,12 +1868,33 @@ dependencies = [ "version_check", ] +[[package]] +name = "unicode-bidi" +version = "0.3.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c1cb5db39152898a79168971543b1cb5020dff7fe43c8dc468b0885f5e29df5" + [[package]] name = "unicode-ident" version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" +[[package]] +name = "unicode-normalization" +version = "0.1.25" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5fd4f6878c9cb28d874b009da9e8d183b5abc80117c40bbd187a1fde336be6e8" +dependencies = [ + "tinyvec", +] + +[[package]] +name = "unicode-properties" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7df058c713841ad818f1dc5d3fd88063241cc61f49f5fbea4b951e8cf5a8d71d" + [[package]] name = "universal-hash" version = "0.5.1" @@ -1459,6 +1913,7 @@ checksum = "144d6b123cef80b301b8f72a9e2ca4370ddec21950d0a103dd22c437006d2db7" dependencies = [ "getrandom 0.4.3", "js-sys", + "serde_core", "wasm-bindgen", ] @@ -1468,6 +1923,12 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65" +[[package]] +name = "vcpkg" +version = "0.2.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426" + [[package]] name = "version_check" version = "0.9.5" @@ -1480,6 +1941,15 @@ version = "0.11.1+wasi-snapshot-preview1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" +[[package]] +name = "wasi" +version = "0.14.7+wasi-0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "883478de20367e224c0090af9cf5f9fa85bed63a95c1abf3afc5c083ebc06e8c" +dependencies = [ + "wasip2", +] + [[package]] name = "wasip2" version = "1.0.4+wasi-0.2.12" @@ -1489,6 +1959,15 @@ dependencies = [ "wit-bindgen", ] +[[package]] +name = "wasite" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "66fe902b4a6b8028a753d5424909b764ccf79b7a209eac9bf97e59cda9f71a42" +dependencies = [ + "wasi 0.14.7+wasi-0.2.4", +] + [[package]] name = "wasm-bindgen" version = "0.2.125" @@ -1534,6 +2013,29 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "web-sys" +version = "0.3.102" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6430a72df5eb332242960fe84b3002a241163998241eb596d4f739b9757061d" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "whoami" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6a5b12f9df4f978d2cfdb1bd3bac52433f44393342d7ee9c25f5a1c14c0f45d" +dependencies = [ + "libc", + "libredox", + "objc2-system-configuration", + "wasite", + "web-sys", +] + [[package]] name = "windows-link" version = "0.2.1" @@ -1549,6 +2051,18 @@ dependencies = [ "windows-link", ] +[[package]] +name = "winnow" +version = "0.7.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945" + +[[package]] +name = "winnow" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0592e1c9d151f854e6fd382574c3a0855250e1d9b2f99d9281c6e6391af352f1" + [[package]] name = "wit-bindgen" version = "0.57.1" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index e154924..f912485 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -12,12 +12,19 @@ axum-plugin = { git = "https://git.fasharp.io/fa-sharp/axum-plugin", rev = "9f72278b3c" } +diesel = { + version = "2.3.10", + default-features = false, + features = ["postgres", "time", "uuid"] +} +diesel-async = { version = "0.9.2", features = ["deadpool", "postgres"] } +diesel_migrations = { version = "2.3.2", features = ["postgres"] } dotenvy = "0.15.7" figment = { version = "0.10.19", features = ["env"] } hex = "0.4.3" serde = { version = "1.0.228", features = ["derive"] } serde_json = "1.0.150" -time = { version = "=0.3.49", default-features = false } +time = { version = "0.3.51", features = ["serde"] } tokio = { version = "1.52.3", default-features = false, @@ -41,3 +48,4 @@ tower-sessions = { tracing = "0.1.44" tracing-appender = "0.2.5" tracing-subscriber = { version = "0.3.23", features = ["env-filter", "json"] } +uuid = { version = "1.23.3", features = ["serde", "v4"] } diff --git a/server-new/migrations/00000000000000_diesel_initial_setup/down.sql b/server-new/migrations/00000000000000_diesel_initial_setup/down.sql new file mode 100644 index 0000000..a9f5260 --- /dev/null +++ b/server-new/migrations/00000000000000_diesel_initial_setup/down.sql @@ -0,0 +1,6 @@ +-- This file was automatically created by Diesel to setup helper functions +-- and other internal bookkeeping. This file is safe to edit, any future +-- changes will be added to existing projects as new migrations. + +DROP FUNCTION IF EXISTS diesel_manage_updated_at(_tbl regclass); +DROP FUNCTION IF EXISTS diesel_set_updated_at(); diff --git a/server-new/migrations/00000000000000_diesel_initial_setup/up.sql b/server-new/migrations/00000000000000_diesel_initial_setup/up.sql new file mode 100644 index 0000000..d68895b --- /dev/null +++ b/server-new/migrations/00000000000000_diesel_initial_setup/up.sql @@ -0,0 +1,36 @@ +-- This file was automatically created by Diesel to setup helper functions +-- and other internal bookkeeping. This file is safe to edit, any future +-- changes will be added to existing projects as new migrations. + + + + +-- Sets up a trigger for the given table to automatically set a column called +-- `updated_at` whenever the row is modified (unless `updated_at` was included +-- in the modified columns) +-- +-- # Example +-- +-- ```sql +-- CREATE TABLE users (id SERIAL PRIMARY KEY, updated_at TIMESTAMP NOT NULL DEFAULT NOW()); +-- +-- SELECT diesel_manage_updated_at('users'); +-- ``` +CREATE OR REPLACE FUNCTION diesel_manage_updated_at(_tbl regclass) RETURNS VOID AS $$ +BEGIN + EXECUTE format('CREATE TRIGGER set_updated_at BEFORE UPDATE ON %s + FOR EACH ROW EXECUTE PROCEDURE diesel_set_updated_at()', _tbl); +END; +$$ LANGUAGE plpgsql; + +CREATE OR REPLACE FUNCTION diesel_set_updated_at() RETURNS trigger AS $$ +BEGIN + IF ( + NEW IS DISTINCT FROM OLD AND + NEW.updated_at IS NOT DISTINCT FROM OLD.updated_at + ) THEN + NEW.updated_at := current_timestamp; + END IF; + RETURN NEW; +END; +$$ LANGUAGE plpgsql; diff --git a/server-new/migrations/2025-06-12-171524_create_chat_sessions/down.sql b/server-new/migrations/2025-06-12-171524_create_chat_sessions/down.sql new file mode 100644 index 0000000..0f76f0b --- /dev/null +++ b/server-new/migrations/2025-06-12-171524_create_chat_sessions/down.sql @@ -0,0 +1,5 @@ +DROP TABLE chat_messages; + +DROP TYPE chat_message_role; + +DROP TABLE chat_sessions; diff --git a/server-new/migrations/2025-06-12-171524_create_chat_sessions/up.sql b/server-new/migrations/2025-06-12-171524_create_chat_sessions/up.sql new file mode 100644 index 0000000..55db20f --- /dev/null +++ b/server-new/migrations/2025-06-12-171524_create_chat_sessions/up.sql @@ -0,0 +1,25 @@ +CREATE TABLE chat_sessions ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid (), + title VARCHAR NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW (), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW () +); + +SELECT + diesel_manage_updated_at ('chat_sessions'); + +CREATE TYPE chat_message_role AS ENUM ('user', 'assistant', 'system'); + +CREATE TABLE chat_messages ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid (), + session_id UUID NOT NULL REFERENCES chat_sessions (id) ON UPDATE CASCADE ON DELETE CASCADE, + role chat_message_role NOT NULL, + content TEXT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW (), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW () +); + +CREATE INDEX chat_messages_session_id_idx ON chat_messages (session_id); + +SELECT + diesel_manage_updated_at ('chat_messages'); diff --git a/server-new/migrations/2025-06-15-031207_add_users/down.sql b/server-new/migrations/2025-06-15-031207_add_users/down.sql new file mode 100644 index 0000000..b896b01 --- /dev/null +++ b/server-new/migrations/2025-06-15-031207_add_users/down.sql @@ -0,0 +1,6 @@ +DROP INDEX chat_sessions_user_id_idx; + +ALTER TABLE chat_sessions +DROP COLUMN user_id; + +DROP TABLE users; diff --git a/server-new/migrations/2025-06-15-031207_add_users/up.sql b/server-new/migrations/2025-06-15-031207_add_users/up.sql new file mode 100644 index 0000000..c580279 --- /dev/null +++ b/server-new/migrations/2025-06-15-031207_add_users/up.sql @@ -0,0 +1,15 @@ +CREATE TABLE users ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid (), + github_id VARCHAR NOT NULL, + name VARCHAR NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW (), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW () +); + +SELECT + diesel_manage_updated_at ('users'); + +ALTER TABLE chat_sessions +ADD COLUMN user_id UUID NOT NULL REFERENCES users (id); + +CREATE INDEX chat_sessions_user_id_idx ON chat_sessions (user_id); diff --git a/server-new/migrations/2025-06-16-035815_add_api_keys/down.sql b/server-new/migrations/2025-06-16-035815_add_api_keys/down.sql new file mode 100644 index 0000000..031e93a --- /dev/null +++ b/server-new/migrations/2025-06-16-035815_add_api_keys/down.sql @@ -0,0 +1,3 @@ +DROP TABLE api_keys; + +DROP TYPE llm_provider; diff --git a/server-new/migrations/2025-06-16-035815_add_api_keys/up.sql b/server-new/migrations/2025-06-16-035815_add_api_keys/up.sql new file mode 100644 index 0000000..b3c1168 --- /dev/null +++ b/server-new/migrations/2025-06-16-035815_add_api_keys/up.sql @@ -0,0 +1,12 @@ +CREATE TYPE llm_provider AS ENUM('anthropic', 'openai', 'ollama', 'deepseek', 'google', 'openrouter'); + +CREATE TABLE api_keys ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + user_id UUID NOT NULL REFERENCES users (id), + provider llm_provider NOT NULL, + ciphertext BYTEA NOT NULL, + nonce BYTEA NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); + +CREATE INDEX api_keys_user_id_idx ON api_keys (user_id); diff --git a/server-new/migrations/2025-06-17-063104_add_message_meta/down.sql b/server-new/migrations/2025-06-17-063104_add_message_meta/down.sql new file mode 100644 index 0000000..9f6c31a --- /dev/null +++ b/server-new/migrations/2025-06-17-063104_add_message_meta/down.sql @@ -0,0 +1,2 @@ +ALTER TABLE chat_messages +DROP COLUMN meta; diff --git a/server-new/migrations/2025-06-17-063104_add_message_meta/up.sql b/server-new/migrations/2025-06-17-063104_add_message_meta/up.sql new file mode 100644 index 0000000..73bc069 --- /dev/null +++ b/server-new/migrations/2025-06-17-063104_add_message_meta/up.sql @@ -0,0 +1,2 @@ +ALTER TABLE chat_messages +ADD COLUMN meta JSONB NOT NULL DEFAULT '{}'; diff --git a/server-new/migrations/2025-06-21-031403_full_text_search/down.sql b/server-new/migrations/2025-06-21-031403_full_text_search/down.sql new file mode 100644 index 0000000..2a59b48 --- /dev/null +++ b/server-new/migrations/2025-06-21-031403_full_text_search/down.sql @@ -0,0 +1,10 @@ +DROP TRIGGER chat_sessions_search_vector_update on chat_sessions; + +DROP TRIGGER chat_messages_search_vector_update on chat_messages; + +DROP FUNCTION chat_sessions_search_vector_update; + +DROP FUNCTION chat_messages_search_vector_update; + +ALTER TABLE chat_messages +DROP COLUMN search_vector; diff --git a/server-new/migrations/2025-06-21-031403_full_text_search/up.sql b/server-new/migrations/2025-06-21-031403_full_text_search/up.sql new file mode 100644 index 0000000..0cbb2f7 --- /dev/null +++ b/server-new/migrations/2025-06-21-031403_full_text_search/up.sql @@ -0,0 +1,53 @@ +ALTER TABLE chat_messages +ADD COLUMN search_vector tsvector NOT NULL DEFAULT ''; + +CREATE INDEX chat_messages_search_vector_idx ON chat_messages USING GIN (search_vector); + +UPDATE chat_messages +SET + search_vector = setweight( + to_tsvector( + 'english', + ( + SELECT + title + FROM + chat_sessions + WHERE + id = session_id + ) + ), + 'A' + ) || setweight(to_tsvector('english', "content"), 'B'); + +CREATE OR REPLACE FUNCTION chat_messages_search_vector_update () RETURNS trigger AS $$ +BEGIN + NEW.search_vector := + setweight(to_tsvector('english', ( + SELECT title FROM chat_sessions WHERE id = NEW.session_id + )), 'A') || setweight(to_tsvector('english', NEW."content"), 'B'); + RETURN NEW; +END; +$$ LANGUAGE plpgsql; + +CREATE OR REPLACE FUNCTION chat_sessions_search_vector_update () RETURNS trigger AS $$ + BEGIN + IF old.title = new.title THEN RETURN NEW; END IF; + UPDATE chat_messages + SET search_vector = + setweight(to_tsvector('english', NEW.title), 'A') || setweight(to_tsvector('english', "content"), 'B') + WHERE session_id = NEW.id; + RETURN NEW; + END; +$$ LANGUAGE plpgsql; + +CREATE TRIGGER chat_messages_search_vector_update BEFORE INSERT +OR +UPDATE ON chat_messages FOR EACH ROW +EXECUTE FUNCTION chat_messages_search_vector_update (); + +CREATE TRIGGER chat_sessions_search_vector_update +AFTER INSERT +OR +UPDATE ON chat_sessions FOR EACH ROW +EXECUTE FUNCTION chat_sessions_search_vector_update (); diff --git a/server-new/migrations/2025-06-23-023453_update_session_updated_at/down.sql b/server-new/migrations/2025-06-23-023453_update_session_updated_at/down.sql new file mode 100644 index 0000000..4d4a89d --- /dev/null +++ b/server-new/migrations/2025-06-23-023453_update_session_updated_at/down.sql @@ -0,0 +1,3 @@ +DROP TRIGGER chat_messages_update_session_updated_at ON chat_messages; + +DROP FUNCTION chat_messages_update_session_updated_at (); diff --git a/server-new/migrations/2025-06-23-023453_update_session_updated_at/up.sql b/server-new/migrations/2025-06-23-023453_update_session_updated_at/up.sql new file mode 100644 index 0000000..3485074 --- /dev/null +++ b/server-new/migrations/2025-06-23-023453_update_session_updated_at/up.sql @@ -0,0 +1,11 @@ +CREATE OR REPLACE FUNCTION chat_messages_update_session_updated_at () RETURNS TRIGGER AS $$ +BEGIN + UPDATE chat_sessions SET updated_at = NOW() WHERE id = NEW.session_id; + RETURN NEW; +END; +$$ LANGUAGE plpgsql; + +CREATE TRIGGER chat_messages_update_session_updated_at BEFORE INSERT +OR +UPDATE ON chat_messages FOR EACH ROW +EXECUTE FUNCTION chat_messages_update_session_updated_at (); diff --git a/server-new/migrations/2025-07-10-165227_add_auth_providers/down.sql b/server-new/migrations/2025-07-10-165227_add_auth_providers/down.sql new file mode 100644 index 0000000..c0098b8 --- /dev/null +++ b/server-new/migrations/2025-07-10-165227_add_auth_providers/down.sql @@ -0,0 +1,16 @@ +-- Drop all unique constraints first +ALTER TABLE users +DROP CONSTRAINT IF EXISTS github_id_unique, +DROP CONSTRAINT IF EXISTS google_id_unique, +DROP CONSTRAINT IF EXISTS discord_id_unique, +DROP CONSTRAINT IF EXISTS oidc_id_unique; + +-- Drop all added columns and restore github_id constraint +ALTER TABLE users +DROP COLUMN avatar_url, +DROP COLUMN oidc_id, +DROP COLUMN discord_id, +DROP COLUMN google_id, +DROP COLUMN sso_username, +ALTER COLUMN github_id +SET NOT NULL; diff --git a/server-new/migrations/2025-07-10-165227_add_auth_providers/up.sql b/server-new/migrations/2025-07-10-165227_add_auth_providers/up.sql new file mode 100644 index 0000000..5b323b1 --- /dev/null +++ b/server-new/migrations/2025-07-10-165227_add_auth_providers/up.sql @@ -0,0 +1,16 @@ +-- Migration: Add support for multiple auth providers and avatars +ALTER TABLE users +ALTER COLUMN github_id +DROP NOT NULL, +ADD COLUMN sso_username TEXT, +ADD COLUMN google_id TEXT, +ADD COLUMN discord_id TEXT, +ADD COLUMN oidc_id TEXT, +ADD COLUMN avatar_url TEXT; + +-- Add unique constraints for all provider IDs +ALTER TABLE users +ADD CONSTRAINT github_id_unique UNIQUE (github_id), +ADD CONSTRAINT google_id_unique UNIQUE (google_id), +ADD CONSTRAINT discord_id_unique UNIQUE (discord_id), +ADD CONSTRAINT oidc_id_unique UNIQUE (oidc_id); diff --git a/server-new/migrations/2025-07-11-012329_add_app_api_keys/down.sql b/server-new/migrations/2025-07-11-012329_add_app_api_keys/down.sql new file mode 100644 index 0000000..5e20619 --- /dev/null +++ b/server-new/migrations/2025-07-11-012329_add_app_api_keys/down.sql @@ -0,0 +1 @@ +DROP TABLE app_api_keys; diff --git a/server-new/migrations/2025-07-11-012329_add_app_api_keys/up.sql b/server-new/migrations/2025-07-11-012329_add_app_api_keys/up.sql new file mode 100644 index 0000000..dd1d398 --- /dev/null +++ b/server-new/migrations/2025-07-11-012329_add_app_api_keys/up.sql @@ -0,0 +1,8 @@ +CREATE TABLE app_api_keys ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + user_id UUID NOT NULL REFERENCES users (id), + name TEXT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); + +CREATE INDEX app_api_keys_user_id_idx ON app_api_keys (user_id); diff --git a/server-new/migrations/2025-07-13-170127_add_tools/down.sql b/server-new/migrations/2025-07-13-170127_add_tools/down.sql new file mode 100644 index 0000000..0e424e7 --- /dev/null +++ b/server-new/migrations/2025-07-13-170127_add_tools/down.sql @@ -0,0 +1,13 @@ +-- Remove 'tool' from chat_message_role +ALTER TYPE chat_message_role +RENAME TO chat_message_role_old; + +CREATE TYPE chat_message_role AS ENUM('user', 'assistant', 'system'); + +ALTER TABLE chat_messages +ALTER COLUMN role TYPE chat_message_role USING role::text::chat_message_role; + +DROP TYPE chat_message_role_old; + +-- Drop tools table +DROP TABLE tools; diff --git a/server-new/migrations/2025-07-13-170127_add_tools/up.sql b/server-new/migrations/2025-07-13-170127_add_tools/up.sql new file mode 100644 index 0000000..b1e70ef --- /dev/null +++ b/server-new/migrations/2025-07-13-170127_add_tools/up.sql @@ -0,0 +1,26 @@ +-- Add tools table +CREATE TABLE tools ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + user_id UUID NOT NULL REFERENCES users (id), + name TEXT NOT NULL, + description TEXT NOT NULL, + config JSONB NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +SELECT + diesel_manage_updated_at ('tools'); + +CREATE INDEX tools_user_id_idx ON tools (user_id); + +-- Add tool role to chat messages +ALTER TYPE chat_message_role +RENAME TO chat_message_role_old; + +CREATE TYPE chat_message_role AS ENUM('user', 'assistant', 'system', 'tool'); + +ALTER TABLE chat_messages +ALTER COLUMN role TYPE chat_message_role USING role::text::chat_message_role; + +DROP TYPE chat_message_role_old; diff --git a/server-new/migrations/2025-07-17-223807_add_providers/down.sql b/server-new/migrations/2025-07-17-223807_add_providers/down.sql new file mode 100644 index 0000000..02d8f8a --- /dev/null +++ b/server-new/migrations/2025-07-17-223807_add_providers/down.sql @@ -0,0 +1,7 @@ +ALTER TABLE secrets +DROP COLUMN name; + +DROP TABLE providers; + +ALTER TABLE secrets +RENAME TO api_keys; diff --git a/server-new/migrations/2025-07-17-223807_add_providers/up.sql b/server-new/migrations/2025-07-17-223807_add_providers/up.sql new file mode 100644 index 0000000..a16799f --- /dev/null +++ b/server-new/migrations/2025-07-17-223807_add_providers/up.sql @@ -0,0 +1,68 @@ +-- Rename API keys table to secrets +ALTER TABLE api_keys +RENAME TO secrets; + +-- Create providers table +CREATE TABLE providers ( + id SERIAL PRIMARY KEY, + name TEXT NOT NULL, + provider_type TEXT NOT NULL, + user_id UUID NOT NULL REFERENCES users (id), + base_url TEXT, + default_model TEXT NOT NULL, + api_key_id UUID REFERENCES secrets (id) ON UPDATE CASCADE ON DELETE SET NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE INDEX idx_providers_user_id ON providers (user_id); + +-- Create providers for users with existing API keys +INSERT INTO + providers (provider_type, name, user_id, base_url, default_model, api_key_id) +SELECT + 'openai', + 'OpenAI', + secrets.user_id, + NULL, + 'gpt-4o-mini', + id +FROM + secrets +WHERE + secrets.provider = 'openai'; + +INSERT INTO + providers (provider_type, name, user_id, base_url, default_model, api_key_id) +SELECT + 'openai', + 'OpenRouter', + secrets.user_id, + 'https://openrouter.ai/api/v1', + 'openai/gpt-4o-mini', + id +FROM + secrets +WHERE + secrets.provider = 'openrouter'; + +INSERT INTO + providers (provider_type, name, user_id, base_url, default_model, api_key_id) +SELECT + 'anthropic', + 'Anthropic', + secrets.user_id, + NULL, + 'claude-3-7-sonnet-latest', + id +FROM + secrets +WHERE + secrets.provider = 'anthropic'; + +-- Add name to secrets table +ALTER TABLE secrets +ADD COLUMN name TEXT NOT NULL DEFAULT 'api_key'; + +UPDATE secrets +SET + name = secrets.provider || '_api_key'; diff --git a/server-new/migrations/2025-08-08-080101_add_session_meta/down.sql b/server-new/migrations/2025-08-08-080101_add_session_meta/down.sql new file mode 100644 index 0000000..b090cff --- /dev/null +++ b/server-new/migrations/2025-08-08-080101_add_session_meta/down.sql @@ -0,0 +1,2 @@ +ALTER TABLE chat_sessions +DROP COLUMN meta; diff --git a/server-new/migrations/2025-08-08-080101_add_session_meta/up.sql b/server-new/migrations/2025-08-08-080101_add_session_meta/up.sql new file mode 100644 index 0000000..7ef1194 --- /dev/null +++ b/server-new/migrations/2025-08-08-080101_add_session_meta/up.sql @@ -0,0 +1,2 @@ +ALTER TABLE chat_sessions +ADD COLUMN meta JSONB NOT NULL DEFAULT '{}'; diff --git a/server-new/migrations/2025-08-09-160537_remove_unused_secret_field/down.sql b/server-new/migrations/2025-08-09-160537_remove_unused_secret_field/down.sql new file mode 100644 index 0000000..4c9c7ea --- /dev/null +++ b/server-new/migrations/2025-08-09-160537_remove_unused_secret_field/down.sql @@ -0,0 +1,4 @@ +CREATE TYPE llm_provider AS ENUM('anthropic', 'openai', 'ollama', 'deepseek', 'google', 'openrouter'); + +ALTER TABLE secrets +ADD COLUMN provider llm_provider NOT NULL DEFAULT 'openai'; diff --git a/server-new/migrations/2025-08-09-160537_remove_unused_secret_field/up.sql b/server-new/migrations/2025-08-09-160537_remove_unused_secret_field/up.sql new file mode 100644 index 0000000..af9e4aa --- /dev/null +++ b/server-new/migrations/2025-08-09-160537_remove_unused_secret_field/up.sql @@ -0,0 +1,4 @@ +ALTER TABLE secrets +DROP COLUMN provider; + +DROP TYPE llm_provider; diff --git a/server-new/migrations/2025-08-16-152113_refactor_tools/down.sql b/server-new/migrations/2025-08-16-152113_refactor_tools/down.sql new file mode 100644 index 0000000..6d82328 --- /dev/null +++ b/server-new/migrations/2025-08-16-152113_refactor_tools/down.sql @@ -0,0 +1,3 @@ +DROP TABLE external_api_tools; + +DROP TABLE system_tools; diff --git a/server-new/migrations/2025-08-16-152113_refactor_tools/up.sql b/server-new/migrations/2025-08-16-152113_refactor_tools/up.sql new file mode 100644 index 0000000..b5bdd71 --- /dev/null +++ b/server-new/migrations/2025-08-16-152113_refactor_tools/up.sql @@ -0,0 +1,23 @@ +CREATE TABLE system_tools ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + user_id UUID NOT NULL REFERENCES users (id), + data JSONB NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); + +SELECT + diesel_manage_updated_at ('system_tools'); + +CREATE TABLE external_api_tools ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + user_id UUID NOT NULL REFERENCES users (id), + data JSONB NOT NULL, + secret_1 UUID REFERENCES secrets (id) ON UPDATE CASCADE ON DELETE SET NULL, + secret_2 UUID REFERENCES secrets (id) ON UPDATE CASCADE ON DELETE SET NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); + +SELECT + diesel_manage_updated_at ('external_api_tools'); diff --git a/server-new/migrations/2025-08-31-034235_add_files/down.sql b/server-new/migrations/2025-08-31-034235_add_files/down.sql new file mode 100644 index 0000000..38a7300 --- /dev/null +++ b/server-new/migrations/2025-08-31-034235_add_files/down.sql @@ -0,0 +1 @@ +DROP TABLE files; diff --git a/server-new/migrations/2025-08-31-034235_add_files/up.sql b/server-new/migrations/2025-08-31-034235_add_files/up.sql new file mode 100644 index 0000000..bbbc781 --- /dev/null +++ b/server-new/migrations/2025-08-31-034235_add_files/up.sql @@ -0,0 +1,18 @@ +CREATE TABLE files ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + user_id UUID NOT NULL REFERENCES users (id), + session_id UUID REFERENCES chat_sessions (id) ON UPDATE CASCADE ON DELETE SET NULL, + path TEXT NOT NULL, + file_type TEXT NOT NULL, + content_type TEXT NOT NULL, + size INTEGER NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +SELECT + diesel_manage_updated_at ('files'); + +CREATE INDEX idx_files_user_id ON files (user_id); + +CREATE INDEX idx_files_session_id ON files (session_id); diff --git a/server-new/migrations/2025-09-03-063406_remove_old_tools/down.sql b/server-new/migrations/2025-09-03-063406_remove_old_tools/down.sql new file mode 100644 index 0000000..dab05c2 --- /dev/null +++ b/server-new/migrations/2025-09-03-063406_remove_old_tools/down.sql @@ -0,0 +1,15 @@ +-- Add back tools table +CREATE TABLE tools ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + user_id UUID NOT NULL REFERENCES users (id), + name TEXT NOT NULL, + description TEXT NOT NULL, + config JSONB NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +SELECT + diesel_manage_updated_at ('tools'); + +CREATE INDEX tools_user_id_idx ON tools (user_id); diff --git a/server-new/migrations/2025-09-03-063406_remove_old_tools/up.sql b/server-new/migrations/2025-09-03-063406_remove_old_tools/up.sql new file mode 100644 index 0000000..f019705 --- /dev/null +++ b/server-new/migrations/2025-09-03-063406_remove_old_tools/up.sql @@ -0,0 +1,2 @@ +-- Drop old tools table +DROP TABLE tools; diff --git a/server-new/src/config.rs b/server-new/src/config.rs index e88390b..2214ad7 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -18,6 +18,11 @@ pub struct AppConfig { pub log_level: String, #[serde(default = "default_request_id_header")] pub request_id_header: String, + pub ip_header: Option, + + // Database + #[serde(default = "default_database_url")] + pub database_url: String, // Auth pub cookie_key: String, @@ -44,13 +49,15 @@ fn default_log_level() -> String { fn default_request_id_header() -> String { "x-request-id".to_string() } +fn default_database_url() -> String { + "postgres://localhost".to_owned() +} fn default_cookie_name() -> String { "auth-rs-chat".to_string() } fn default_session_length() -> i64 { 60 * 60 * 24 * 7 // 1 week } - fn default_body_limit() -> usize { 2 * 1024 * 1024 // 2 MB } diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs new file mode 100644 index 0000000..26de85c --- /dev/null +++ b/server-new/src/db/mod.rs @@ -0,0 +1,34 @@ +use anyhow::Context; + +use crate::error::AppResult; + +pub mod models; +mod repositories; +mod schema; + +/// Type of the database pool +pub type DbPool = diesel_async::pooled_connection::deadpool::Pool; +/// Type of the database connection retrieved from the pool +pub type DbConnection = + diesel_async::pooled_connection::deadpool::Object; + +/// Wrapper around a database connection that gives access to the repositories, +/// e.g. `UserRepository`, `ChatRepository`, etc. +pub struct DbService { + cxn: DbConnection, +} + +impl DbService { + pub fn new(cxn: DbConnection) -> Self { + Self { cxn } + } + + pub async fn from_pool(pool: &DbPool) -> AppResult { + let cxn = pool.get().await.context("error retrieving DB connection")?; + Ok(Self::new(cxn)) + } + + pub fn users(&mut self) -> repositories::UserRepository<'_> { + repositories::UserRepository::new(&mut self.cxn) + } +} diff --git a/server-new/src/db/models.rs b/server-new/src/db/models.rs new file mode 100644 index 0000000..10430d1 --- /dev/null +++ b/server-new/src/db/models.rs @@ -0,0 +1,17 @@ +use crate::db::schema; + +// mod api_key; +// mod chat; +// mod file; +// mod provider; +// mod secret; +// mod tool; +mod user; + +// pub use api_key::*; +// pub use chat::*; +// pub use file::*; +// pub use provider::*; +// pub use secret::*; +// pub use tool::*; +pub use user::*; diff --git a/server-new/src/db/models/user.rs b/server-new/src/db/models/user.rs new file mode 100644 index 0000000..4383f7b --- /dev/null +++ b/server-new/src/db/models/user.rs @@ -0,0 +1,49 @@ +use diesel::prelude::*; +use serde::Serialize; +use time::OffsetDateTime; +use uuid::Uuid; + +#[derive(Identifiable, Queryable, Selectable, Serialize)] +#[diesel(table_name = super::schema::users)] +pub struct ChatRsUser { + pub id: Uuid, + pub name: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub avatar_url: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub github_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub google_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub discord_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub oidc_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub sso_username: Option, + pub created_at: OffsetDateTime, + pub updated_at: OffsetDateTime, +} + +#[derive(Insertable, Default)] +#[diesel(table_name = super::schema::users)] +pub struct NewChatRsUser<'r> { + pub github_id: Option<&'r str>, + pub google_id: Option<&'r str>, + pub discord_id: Option<&'r str>, + pub oidc_id: Option<&'r str>, + pub sso_username: Option<&'r str>, + pub name: &'r str, + pub avatar_url: Option<&'r str>, +} + +#[derive(AsChangeset, Default)] +#[diesel(table_name = super::schema::users)] +pub struct UpdateChatRsUser<'r> { + pub github_id: Option<&'r str>, + pub google_id: Option<&'r str>, + pub discord_id: Option<&'r str>, + pub oidc_id: Option<&'r str>, + pub sso_username: Option<&'r str>, + pub name: Option<&'r str>, + pub avatar_url: Option<&'r str>, +} diff --git a/server-new/src/db/repositories/mod.rs b/server-new/src/db/repositories/mod.rs new file mode 100644 index 0000000..5ba8673 --- /dev/null +++ b/server-new/src/db/repositories/mod.rs @@ -0,0 +1,3 @@ +mod user; + +pub use user::UserRepository; diff --git a/server-new/src/db/repositories/user.rs b/server-new/src/db/repositories/user.rs new file mode 100644 index 0000000..9e211e4 --- /dev/null +++ b/server-new/src/db/repositories/user.rs @@ -0,0 +1,117 @@ +use diesel::prelude::*; +use diesel::result::Error; +use diesel_async::RunQueryDsl; +use uuid::Uuid; + +use crate::db::{ + DbConnection, + models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, + schema::users, +}; + +pub struct UserRepository<'a> { + db: &'a mut DbConnection, +} + +impl<'a> UserRepository<'a> { + pub fn new(db: &'a mut DbConnection) -> Self { + UserRepository { db } + } + + pub async fn find_by_id(&mut self, id: &Uuid) -> Result, Error> { + let user = users::table + .filter(users::id.eq(id)) + .select(ChatRsUser::as_select()) + .first(self.db) + .await + .optional()?; + + Ok(user) + } + + // pub async fn find_by_github_id(&mut self, id: &str) -> Result, Error> { + // let user = users::table + // .filter(users::github_id.eq(id)) + // .select(ChatRsUser::as_select()) + // .first(self.db) + // .await + // .optional()?; + + // Ok(user) + // } + + // pub async fn find_by_google_id(&mut self, id: &str) -> Result, Error> { + // let user = users::table + // .filter(users::google_id.eq(id)) + // .select(ChatRsUser::as_select()) + // .first(self.db) + // .await + // .optional()?; + + // Ok(user) + // } + + // pub async fn find_by_discord_id(&mut self, id: &str) -> Result, Error> { + // let user = users::table + // .filter(users::discord_id.eq(id)) + // .select(ChatRsUser::as_select()) + // .first(self.db) + // .await + // .optional()?; + + // Ok(user) + // } + + // pub async fn find_by_oidc_id(&mut self, id: &str) -> Result, Error> { + // let user = users::table + // .filter(users::oidc_id.eq(id)) + // .select(ChatRsUser::as_select()) + // .first(self.db) + // .await + // .optional()?; + + // Ok(user) + // } + + // pub async fn find_by_sso_username(&mut self, username: &str) -> Result, Error> { + // let user_id = users::table + // .filter(users::sso_username.eq(username)) + // .select(users::id) + // .first(self.db) + // .await + // .optional()?; + + // Ok(user_id) + // } + + // pub async fn create(&mut self, user: NewChatRsUser<'_>) -> Result { + // diesel::insert_into(users::table) + // .values(user) + // .returning(ChatRsUser::as_returning()) + // .get_result(self.db) + // .await + // } + + // pub async fn update( + // &mut self, + // user_id: &Uuid, + // data: UpdateChatRsUser<'_>, + // ) -> Result { + // let updated_id: Uuid = diesel::update(users::table.find(user_id)) + // .set(data) + // .returning(users::id) + // .get_result(self.db) + // .await?; + + // Ok(updated_id) + // } + + // pub async fn delete(&mut self, user_id: &Uuid) -> Result { + // let id: Uuid = diesel::delete(users::table.find(user_id)) + // .returning(users::id) + // .get_result(self.db) + // .await?; + + // Ok(id) + // } +} diff --git a/server-new/src/db/schema.rs b/server-new/src/db/schema.rs new file mode 100644 index 0000000..2485dab --- /dev/null +++ b/server-new/src/db/schema.rs @@ -0,0 +1,146 @@ +// @generated automatically by Diesel CLI. + +pub mod sql_types { + #[derive(diesel::query_builder::QueryId, Clone, diesel::sql_types::SqlType)] + #[diesel(postgres_type(name = "chat_message_role"))] + pub struct ChatMessageRole; + + #[derive(diesel::query_builder::QueryId, Clone, diesel::sql_types::SqlType)] + #[diesel(postgres_type(name = "tsvector", schema = "pg_catalog"))] + pub struct Tsvector; +} + +diesel::table! { + app_api_keys (id) { + id -> Uuid, + user_id -> Uuid, + name -> Text, + created_at -> Timestamptz, + } +} + +diesel::table! { + use diesel::sql_types::*; + use super::sql_types::ChatMessageRole; + use super::sql_types::Tsvector; + + chat_messages (id) { + id -> Uuid, + session_id -> Uuid, + role -> ChatMessageRole, + content -> Text, + created_at -> Timestamptz, + updated_at -> Timestamptz, + meta -> Jsonb, + search_vector -> Tsvector, + } +} + +diesel::table! { + chat_sessions (id) { + id -> Uuid, + title -> Varchar, + created_at -> Timestamptz, + updated_at -> Timestamptz, + user_id -> Uuid, + meta -> Jsonb, + } +} + +diesel::table! { + external_api_tools (id) { + id -> Uuid, + user_id -> Uuid, + data -> Jsonb, + secret_1 -> Nullable, + secret_2 -> Nullable, + created_at -> Timestamptz, + updated_at -> Timestamptz, + } +} + +diesel::table! { + files (id) { + id -> Uuid, + user_id -> Uuid, + session_id -> Nullable, + path -> Text, + file_type -> Text, + content_type -> Text, + size -> Int4, + created_at -> Timestamptz, + updated_at -> Timestamptz, + } +} + +diesel::table! { + providers (id) { + id -> Int4, + name -> Text, + provider_type -> Text, + user_id -> Uuid, + base_url -> Nullable, + default_model -> Text, + api_key_id -> Nullable, + created_at -> Timestamptz, + } +} + +diesel::table! { + secrets (id) { + id -> Uuid, + user_id -> Uuid, + ciphertext -> Bytea, + nonce -> Bytea, + created_at -> Timestamptz, + name -> Text, + } +} + +diesel::table! { + system_tools (id) { + id -> Uuid, + user_id -> Uuid, + data -> Jsonb, + created_at -> Timestamptz, + updated_at -> Timestamptz, + } +} + +diesel::table! { + users (id) { + id -> Uuid, + github_id -> Nullable, + name -> Varchar, + created_at -> Timestamptz, + updated_at -> Timestamptz, + sso_username -> Nullable, + google_id -> Nullable, + discord_id -> Nullable, + oidc_id -> Nullable, + avatar_url -> Nullable, + } +} + +diesel::joinable!(app_api_keys -> users (user_id)); +diesel::joinable!(chat_messages -> chat_sessions (session_id)); +diesel::joinable!(chat_sessions -> users (user_id)); +diesel::joinable!(external_api_tools -> users (user_id)); +diesel::joinable!(files -> chat_sessions (session_id)); +diesel::joinable!(files -> users (user_id)); +diesel::joinable!(providers -> secrets (api_key_id)); +diesel::joinable!(providers -> users (user_id)); +diesel::joinable!(secrets -> users (user_id)); +diesel::joinable!(system_tools -> users (user_id)); + +diesel::allow_tables_to_appear_in_same_query!( + app_api_keys, + chat_messages, + chat_sessions, + external_api_tools, + files, + providers, + secrets, + system_tools, + users, +); diff --git a/server-new/src/error.rs b/server-new/src/error.rs index c0b2e0f..b75e3f3 100644 --- a/server-new/src/error.rs +++ b/server-new/src/error.rs @@ -22,12 +22,11 @@ impl From for AppError { } } -// Add more conversions here to be able to propagate them in route handlers: -// impl From for AppError { -// fn from(error: DatabaseError) -> Self { -// Self::internal(error.into()) -// } -// } +impl From for AppError { + fn from(error: diesel::result::Error) -> Self { + Self::internal(error.into()) + } +} impl AppError { pub fn new(status: StatusCode, message: impl Into) -> Self { diff --git a/server-new/src/extractors/session.rs b/server-new/src/extractors/session.rs index 4ced4aa..c55271e 100644 --- a/server-new/src/extractors/session.rs +++ b/server-new/src/extractors/session.rs @@ -1,12 +1,24 @@ +use std::{ + net::{IpAddr, SocketAddr}, + str::FromStr, +}; + use anyhow::Context; -use axum::extract::{FromRequestParts, OptionalFromRequestParts}; +use axum::{ + extract::{ConnectInfo, FromRequestParts, OptionalFromRequestParts}, + http::header, +}; use serde::{Deserialize, Serialize}; +use time::UtcDateTime; use tower_sessions::Session; +use uuid::Uuid; -use crate::error::AppError; +use crate::{error::AppError, state::AppState}; -/// The field in the `tower_sessions` storage used to store the user session data +/// The field used to store the user session data const USER_SESSION_FIELD: &str = "sess"; +/// The field used to store the user session metadata +const SESSION_META_FIELD: &str = "meta"; /// Active session data. Beware when changing or adding to this struct, as it can /// invalidate existing sessions. @@ -18,19 +30,34 @@ const USER_SESSION_FIELD: &str = "sess"; /// if there is no active session. #[derive(Debug, Serialize, Deserialize)] pub struct UserSession { - pub user_id: String, + pub user_id: Uuid, +} + +/// Session metadata extracted on login. +#[derive(Debug, Serialize, Deserialize)] +pub struct SessionMeta { + pub start_time: UtcDateTime, + pub ip: Option, + pub user_agent: Option, } impl UserSession { - pub fn new(user_id: String) -> Self { + pub fn new(user_id: Uuid) -> Self { Self { user_id } } - pub async fn init(session: &Session, user_id: &str) -> Result<(), AppError> { - Ok(session - .insert(USER_SESSION_FIELD, UserSession::new(user_id.into())) + pub async fn init( + session: &Session, + meta: &SessionMeta, + user_id: &Uuid, + ) -> Result<(), AppError> { + session + .insert(USER_SESSION_FIELD, UserSession::new(user_id.clone())) .await - .context("failed to initialize session")?) + .context("failed to initialize session")?; + session.insert(SESSION_META_FIELD, meta).await.ok(); + + Ok(()) } } @@ -68,3 +95,38 @@ impl FromRequestParts for UserSession { } } } + +impl FromRequestParts for SessionMeta { + type Rejection = (); + + async fn from_request_parts( + parts: &mut axum::http::request::Parts, + state: &AppState, + ) -> Result { + let ip_header = state + .config + .ip_header + .as_ref() + .and_then(|h| parts.headers.get(h).and_then(|h| h.to_str().ok())) + .and_then(|h| IpAddr::from_str(h).ok()); + let ip = match ip_header { + Some(ip) => Some(ip), + None => ConnectInfo::::from_request_parts(parts, state) + .await + .ok() + .map(|info| info.ip()), + }; + let user_agent = parts + .headers + .get(header::USER_AGENT) + .and_then(|h| h.to_str().ok()) + .map(|ua| ua.to_owned()); + let start_time = UtcDateTime::now(); + + Ok(Self { + start_time, + ip, + user_agent, + }) + } +} diff --git a/server-new/src/lib.rs b/server-new/src/lib.rs index 7de7e9b..38ae43f 100644 --- a/server-new/src/lib.rs +++ b/server-new/src/lib.rs @@ -3,15 +3,18 @@ use axum_plugin::{App, InitializedApp}; use crate::state::AppState; mod config; +mod db; mod error; mod extractors; mod plugins; mod routes; +mod services; mod state; pub async fn create_app() -> anyhow::Result> { let app = App::new() .register(config::plugin()) // Extract configuration and add to state + .register(plugins::database::plugin()) // Initialize database .register(routes::plugin()) // Add API routes .register(plugins::session::plugin()) // Setup sessions .register(plugins::logging::plugin()) // Request logging diff --git a/server-new/src/main.rs b/server-new/src/main.rs index 2907e3e..655290e 100644 --- a/server-new/src/main.rs +++ b/server-new/src/main.rs @@ -27,9 +27,13 @@ async fn main() -> anyhow::Result<()> { let addr = SocketAddr::new(config.host, config.port); let listener = tokio::net::TcpListener::bind(addr).await?; tracing::info!("Server listening on http://{}...", listener.local_addr()?); - axum::serve(listener, app.router().into_make_service()) - .with_graceful_shutdown(shutdown_signal(app.shutdown())) - .await?; + axum::serve( + listener, + app.router() + .into_make_service_with_connect_info::(), + ) + .with_graceful_shutdown(shutdown_signal(app.shutdown())) + .await?; Ok(()) } @@ -38,11 +42,7 @@ fn init_logging() -> ( tracing_subscriber::reload::Handle, tracing_appender::non_blocking::WorkerGuard, ) { - let init_log_level = if cfg!(debug_assertions) { - dotenvy::var("RS_CHAT_LOG_LEVEL").unwrap_or("info".into()) - } else { - "info".to_owned() - }; + let init_log_level = std::env::var("RS_CHAT_LOG_LEVEL").unwrap_or("info".into()); let (writer, guard) = tracing_appender::non_blocking(std::io::stdout()); let (filter_layer, filter_handle) = tracing_subscriber::reload::Layer::new(EnvFilter::new(init_log_level)); diff --git a/server-new/src/plugins/database.rs b/server-new/src/plugins/database.rs new file mode 100644 index 0000000..58c32ad --- /dev/null +++ b/server-new/src/plugins/database.rs @@ -0,0 +1,51 @@ +use anyhow::Context; +use axum_plugin::AdHocPlugin; +use diesel::Connection; +use diesel_async::{ + AsyncPgConnection, + pooled_connection::{AsyncDieselConnectionManager, deadpool::Pool}, +}; +use diesel_migrations::{EmbeddedMigrations, MigrationHarness}; + +use crate::{config::AppConfig, db::DbPool, state::AppState}; + +const MIGRATIONS: EmbeddedMigrations = diesel_migrations::embed_migrations!(); + +pub fn plugin() -> AdHocPlugin { + AdHocPlugin::named("Database") + .on_init(async |mut state| { + let app_config = state.get::().context("missing config")?; + let database_url = app_config.database_url.clone(); + + tokio::task::spawn_blocking(move || { + let mut cxn = diesel::PgConnection::establish(&database_url) + .context("Failed to connect to database")?; + tracing::info!("Connected to database at '{database_url}'"); + match cxn.run_pending_migrations(MIGRATIONS) { + Ok(run_migrations) => { + for migration in run_migrations { + tracing::info!("Migration run: '{migration}'"); + } + Ok(()) + } + Err(err) => Err(anyhow::anyhow!(err.to_string()).context("Migration failed")), + } + }) + .await??; + + let manager = + AsyncDieselConnectionManager::::new(&app_config.database_url); + let pool: DbPool = Pool::builder(manager).build()?; + + state.insert(pool); + Ok(state) + }) + .on_shutdown(|state: &AppState| { + let pool = state.db_pool.clone(); + async move { + pool.close(); + tracing::info!("Shut down database pool"); + Ok(()) + } + }) +} diff --git a/server-new/src/plugins/logging.rs b/server-new/src/plugins/logging.rs index bdd85d2..f9fc373 100644 --- a/server-new/src/plugins/logging.rs +++ b/server-new/src/plugins/logging.rs @@ -8,16 +8,12 @@ use tower_http::{ request_id::{MakeRequestUuid, PropagateRequestIdLayer, SetRequestIdLayer}, trace::{DefaultOnRequest, DefaultOnResponse, TraceLayer}, }; -use tracing::{Level, level_filters::LevelFilter}; +use tracing::Level; use crate::state::AppState; pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Request logs").on_setup(|router, state: &AppState| { - if LevelFilter::from_str(&state.config.log_level)? > LevelFilter::INFO { - return Ok(router); - } - const LOG_LEVEL: Level = Level::INFO; let request_id_header = HeaderName::from_str(&state.config.request_id_header) .context("invalid request ID header")?; diff --git a/server-new/src/plugins/mod.rs b/server-new/src/plugins/mod.rs index e40cded..e439465 100644 --- a/server-new/src/plugins/mod.rs +++ b/server-new/src/plugins/mod.rs @@ -1,3 +1,4 @@ +pub mod database; pub mod logging; pub mod security; pub mod session; diff --git a/server-new/src/routes/auth.rs b/server-new/src/routes/auth.rs index defeef4..90acbd8 100644 --- a/server-new/src/routes/auth.rs +++ b/server-new/src/routes/auth.rs @@ -1,7 +1,11 @@ use anyhow::Context; -use axum::response::IntoResponse; +use axum::{extract::State, response::IntoResponse}; -use crate::{error::AppResult, extractors::session::UserSession, state::AppState}; +use crate::{ + error::{AppError, AppResult}, + extractors::session::{SessionMeta, UserSession}, + state::AppState, +}; pub fn routes() -> axum::Router { axum::Router::new() @@ -10,16 +14,29 @@ pub fn routes() -> axum::Router { .route("/logout", axum::routing::post(logout_handler)) } -async fn login_handler(session: tower_sessions::Session) -> AppResult { +async fn login_handler( + maybe_user: Option, + session: tower_sessions::Session, + meta: SessionMeta, +) -> AppResult { + if maybe_user.is_some() { + return Err(AppError::bad_request("already logged in")); + } + // TODO login handling logic - let user_id = "user123"; - UserSession::init(&session, user_id).await?; + let user_id = uuid::Uuid::new_v4(); + UserSession::init(&session, &meta, &user_id).await?; Ok(format!("Logged in as {user_id}")) } -async fn get_user_handler(UserSession { user_id }: UserSession) -> impl IntoResponse { - format!("Logged in as {user_id}") +async fn get_user_handler( + UserSession { user_id }: UserSession, + State(state): State, +) -> AppResult { + state.user_service().get_user(&user_id).await?; + + Ok(format!("Logged in as {user_id}")) } async fn logout_handler( diff --git a/server-new/src/services/mod.rs b/server-new/src/services/mod.rs new file mode 100644 index 0000000..cbbc74c --- /dev/null +++ b/server-new/src/services/mod.rs @@ -0,0 +1,3 @@ +mod user; + +pub use user::UserService; diff --git a/server-new/src/services/user.rs b/server-new/src/services/user.rs new file mode 100644 index 0000000..118b380 --- /dev/null +++ b/server-new/src/services/user.rs @@ -0,0 +1,23 @@ +use uuid::Uuid; + +use crate::{ + db::{DbPool, DbService, models::ChatRsUser}, + error::AppResult, +}; + +pub struct UserService<'a> { + db: &'a DbPool, +} + +impl<'a> UserService<'a> { + pub fn new(db: &'a DbPool) -> Self { + Self { db } + } + + pub async fn get_user(&self, id: &Uuid) -> AppResult> { + let mut db = DbService::from_pool(&self.db).await?; + let user = db.users().find_by_id(id).await?; + + Ok(user) + } +} diff --git a/server-new/src/state.rs b/server-new/src/state.rs index 6b9a0e8..1e87b99 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -4,7 +4,7 @@ use std::{ops::Deref, sync::Arc}; use axum_plugin::{AppState, TypeMap}; -use crate::config::AppConfig; +use crate::{config::AppConfig, db::DbPool, services::UserService}; /// App state stored in the Axum router #[derive(Clone)] @@ -13,7 +13,13 @@ pub struct AppState(Arc); #[derive(AppState)] pub struct AppStateInner { pub config: AppConfig, - // add state here... + pub db_pool: DbPool, +} + +impl AppState { + pub fn user_service(&self) -> UserService<'_> { + UserService::new(&self.db_pool) + } } impl Deref for AppState { From 4754b5ca1845eafe3f24834f7e728a36d135a207 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 22 Jun 2026 23:27:11 -0400 Subject: [PATCH 023/111] service layout, login and logout --- server-new/.env.example | 1 + server-new/Cargo.lock | 237 +++++++++++++++--- server-new/Cargo.toml | 19 +- server-new/diesel.toml | 6 + .../down.sql | 1 + .../up.sql | 11 + server-new/src/db/mod.rs | 17 +- server-new/src/db/models.rs | 2 + server-new/src/db/models/session.rs | 65 +++++ server-new/src/db/models/user.rs | 15 +- server-new/src/db/repositories/mod.rs | 2 + server-new/src/db/repositories/session.rs | 76 ++++++ server-new/src/db/repositories/user.rs | 6 +- server-new/src/db/schema.rs | 13 + server-new/src/error.rs | 16 +- server-new/src/extractors/session.rs | 48 +--- server-new/src/plugins/database.rs | 33 ++- server-new/src/plugins/session.rs | 5 +- server-new/src/routes/auth.rs | 26 +- server-new/src/services/auth/mod.rs | 68 +++++ server-new/src/services/auth/session_store.rs | 117 +++++++++ server-new/src/services/mod.rs | 4 +- server-new/src/services/user.rs | 23 -- server-new/src/state.rs | 6 +- 24 files changed, 660 insertions(+), 157 deletions(-) create mode 100644 server-new/diesel.toml create mode 100644 server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/down.sql create mode 100644 server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/up.sql create mode 100644 server-new/src/db/models/session.rs create mode 100644 server-new/src/db/repositories/session.rs create mode 100644 server-new/src/services/auth/mod.rs create mode 100644 server-new/src/services/auth/session_store.rs delete mode 100644 server-new/src/services/user.rs diff --git a/server-new/.env.example b/server-new/.env.example index ade296f..8c718da 100644 --- a/server-new/.env.example +++ b/server-new/.env.example @@ -4,6 +4,7 @@ RS_CHAT_PORT=8080 # Database RS_CHAT_DATABASE_URL=postgres://postgres:postgres@localhost/postgres +DATABASE_URL=postgres://postgres:postgres@localhost/postgres # Auth RS_CHAT_COOKIE_KEY= # hex secret >=32 bytes, e.g. `openssl rand --hex 32` diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index b909299..bf283cc 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -46,6 +46,15 @@ dependencies = [ "memchr", ] +[[package]] +name = "android_system_properties" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "819e7219dbd41043ac279b19830f2efc897156490d7fd6ea916720117ee66311" +dependencies = [ + "libc", +] + [[package]] name = "anyhow" version = "1.0.102" @@ -78,6 +87,12 @@ version = "1.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0" +[[package]] +name = "autocfg" +version = "1.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" + [[package]] name = "axum" version = "0.8.9" @@ -210,12 +225,34 @@ version = "1.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8ae3f5d315924270530207e2a68396c3cc547f6dca3fbdca317cfb1a51edb593" +[[package]] +name = "cc" +version = "1.2.65" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e228eec9be7c17ccb640b59b36a5cd805ea2a564a4c5e162c2f659fea30d3b96" +dependencies = [ + "find-msvc-tools", + "shlex", +] + [[package]] name = "cfg-if" version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" +[[package]] +name = "chrono" +version = "0.4.45" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1aa79e62e7697b8e29b513a68abacf485adcd1fe8284a4316c5ae868e6633327" +dependencies = [ + "iana-time-zone", + "num-traits", + "serde", + "windows-link", +] + [[package]] name = "cipher" version = "0.4.4" @@ -244,6 +281,12 @@ dependencies = [ "version_check", ] +[[package]] +name = "core-foundation-sys" +version = "0.8.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "773648b94d0e5d620f64f280777445740e61fe701025087ec8b57f45c791888b" + [[package]] name = "cpufeatures" version = "0.2.17" @@ -294,8 +337,18 @@ version = "0.21.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9cdf337090841a411e2a7f3deb9187445851f91b309c0c0a29e05f74a00a48c0" dependencies = [ - "darling_core", - "darling_macro", + "darling_core 0.21.3", + "darling_macro 0.21.3", +] + +[[package]] +name = "darling" +version = "0.23.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "25ae13da2f202d56bd7f91c25fba009e7717a1e4a1cc98a76d844b65ae912e9d" +dependencies = [ + "darling_core 0.23.0", + "darling_macro 0.23.0", ] [[package]] @@ -312,13 +365,37 @@ dependencies = [ "syn", ] +[[package]] +name = "darling_core" +version = "0.23.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9865a50f7c335f53564bb694ef660825eb8610e0a53d3e11bf1b0d3df31e03b0" +dependencies = [ + "ident_case", + "proc-macro2", + "quote", + "strsim", + "syn", +] + [[package]] name = "darling_macro" version = "0.21.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d38308df82d1080de0afee5d069fa14b0326a88c14f15c5ccda35b4a6c414c81" dependencies = [ - "darling_core", + "darling_core 0.21.3", + "quote", + "syn", +] + +[[package]] +name = "darling_macro" +version = "0.23.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac3984ec7bd6cfa798e62b4a642426a5be0e68f9401cfc2a01e3fa9ea2fcdb8d" +dependencies = [ + "darling_core 0.23.0", "quote", "syn", ] @@ -357,11 +434,11 @@ checksum = "29fe29a87fb84c631ffb3ba21798c4b1f3a964701ba78f0dce4bf8668562ec88" dependencies = [ "bitflags", "byteorder", + "chrono", "diesel_derives", "downcast-rs", "itoa", - "pq-sys", - "time", + "serde_json", "uuid", ] @@ -373,6 +450,7 @@ checksum = "dd39af30158d444884f166fe4c58f35dc40ad71ad017bb59408a3448526ff4bd" dependencies = [ "deadpool", "diesel", + "diesel_migrations", "futures-core", "futures-util", "pin-project-lite", @@ -442,7 +520,7 @@ version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "dd122633e4bef06db27737f21d3738fb89c8f6d5360d6d9d7635dda142a7757e" dependencies = [ - "darling", + "darling 0.21.3", "either", "heck", "proc-macro2", @@ -485,6 +563,12 @@ dependencies = [ "version_check", ] +[[package]] +name = "find-msvc-tools" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" + [[package]] name = "fnv" version = "1.0.7" @@ -764,6 +848,30 @@ dependencies = [ "tower-service", ] +[[package]] +name = "iana-time-zone" +version = "0.1.65" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e31bc9ad994ba00e440a8aa5c9ef0ec67d5cb5e5cb0cc7f8b744a35b389cc470" +dependencies = [ + "android_system_properties", + "core-foundation-sys", + "iana-time-zone-haiku", + "js-sys", + "log", + "wasm-bindgen", + "windows-core", +] + +[[package]] +name = "iana-time-zone-haiku" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f31827a206f56af32e590ba56d5d2d085f558508192593743f16b2306495269f" +dependencies = [ + "cc", +] + [[package]] name = "ident_case" version = "1.0.1" @@ -922,6 +1030,15 @@ version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "521739c6d2bac4aa25192232afe6841231376b2b26d4d9fae5ecf8ca5772e441" +[[package]] +name = "num-traits" +version = "0.2.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" +dependencies = [ + "autocfg", +] + [[package]] name = "num_cpus" version = "1.17.0" @@ -1039,12 +1156,6 @@ version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" -[[package]] -name = "pkg-config" -version = "0.3.33" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" - [[package]] name = "polyval" version = "0.6.2" @@ -1101,17 +1212,6 @@ dependencies = [ "zerocopy", ] -[[package]] -name = "pq-sys" -version = "0.7.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "574ddd6a267294433f140b02a726b0640c43cf7c6f717084684aaa3b285aba61" -dependencies = [ - "libc", - "pkg-config", - "vcpkg", -] - [[package]] name = "proc-macro2" version = "1.0.106" @@ -1245,9 +1345,11 @@ name = "rs-chat-api" version = "0.1.0" dependencies = [ "anyhow", + "async-trait", "axum", "axum-helmet", "axum-plugin", + "chrono", "diesel", "diesel-async", "diesel_migrations", @@ -1256,7 +1358,7 @@ dependencies = [ "hex", "serde", "serde_json", - "time", + "serde_with", "tokio", "tower", "tower-http", @@ -1366,6 +1468,28 @@ dependencies = [ "serde", ] +[[package]] +name = "serde_with" +version = "3.21.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76a5c54c7310e7b8b9577c286d7e399ddd876c3e12b3ed917a8aabc4b96e9e8c" +dependencies = [ + "serde_core", + "serde_with_macros", +] + +[[package]] +name = "serde_with_macros" +version = "3.21.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "84d57bc0c8b9a17920c178daa6bb924850d54a9c97ab45194bb8c17ad66bb660" +dependencies = [ + "darling 0.23.0", + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "sha2" version = "0.10.9" @@ -1386,6 +1510,12 @@ dependencies = [ "lazy_static", ] +[[package]] +name = "shlex" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8fadd59c855ef2080decdef8ff161eb6661b86933c9d82e5ba29dc602a55aba" + [[package]] name = "signal-hook-registry" version = "1.4.8" @@ -1923,12 +2053,6 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65" -[[package]] -name = "vcpkg" -version = "0.2.15" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426" - [[package]] name = "version_check" version = "0.9.5" @@ -2036,12 +2160,65 @@ dependencies = [ "web-sys", ] +[[package]] +name = "windows-core" +version = "0.62.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8e83a14d34d0623b51dce9581199302a221863196a1dde71a7663a4c2be9deb" +dependencies = [ + "windows-implement", + "windows-interface", + "windows-link", + "windows-result", + "windows-strings", +] + +[[package]] +name = "windows-implement" +version = "0.60.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "053e2e040ab57b9dc951b72c264860db7eb3b0200ba345b4e4c3b14f67855ddf" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "windows-interface" +version = "0.59.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f316c4a2570ba26bbec722032c4099d8c8bc095efccdc15688708623367e358" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "windows-link" version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" +[[package]] +name = "windows-result" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7781fa89eaf60850ac3d2da7af8e5242a5ea78d1a11c49bf2910bb5a73853eb5" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-strings" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7837d08f69c77cf6b07689544538e017c1bfcf57e34b4c0ff58e6c2cd3b37091" +dependencies = [ + "windows-link", +] + [[package]] name = "windows-sys" version = "0.61.2" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index f912485..72e75a6 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -6,25 +6,38 @@ description = "LLM chat application" [dependencies] anyhow = "1.0.102" +async-trait = "0.1.89" axum = { version = "0.8.9", features = ["json", "query"] } axum-helmet = "1.0.2" axum-plugin = { git = "https://git.fasharp.io/fa-sharp/axum-plugin", rev = "9f72278b3c" } +chrono = { + version = "0.4.45", + default-features = false, + features = ["now", "serde", "std"] +} diesel = { version = "2.3.10", default-features = false, - features = ["postgres", "time", "uuid"] + features = ["chrono", "serde_json", "uuid"] +} +diesel-async = { + version = "0.9.2", + features = ["deadpool", "migrations", "postgres"] } -diesel-async = { version = "0.9.2", features = ["deadpool", "postgres"] } diesel_migrations = { version = "2.3.2", features = ["postgres"] } dotenvy = "0.15.7" figment = { version = "0.10.19", features = ["env"] } hex = "0.4.3" serde = { version = "1.0.228", features = ["derive"] } serde_json = "1.0.150" -time = { version = "0.3.51", features = ["serde"] } +serde_with = { + version = "3.21.0", + default-features = false, + features = ["macros"] +} tokio = { version = "1.52.3", default-features = false, diff --git a/server-new/diesel.toml b/server-new/diesel.toml new file mode 100644 index 0000000..ad7fef6 --- /dev/null +++ b/server-new/diesel.toml @@ -0,0 +1,6 @@ +# For documentation on how to configure this file, +# see https://diesel.rs/guides/configuring-diesel-cli + +[print_schema] +file = "src/db/schema.rs" +custom_type_derives = ["diesel::query_builder::QueryId", "Clone"] diff --git a/server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/down.sql b/server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/down.sql new file mode 100644 index 0000000..25b7a7f --- /dev/null +++ b/server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/down.sql @@ -0,0 +1 @@ +DROP TABLE auth_sessions; diff --git a/server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/up.sql b/server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/up.sql new file mode 100644 index 0000000..3c7c368 --- /dev/null +++ b/server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/up.sql @@ -0,0 +1,11 @@ +CREATE TABLE auth_sessions ( + id UUID PRIMARY KEY, + user_id UUID NOT NULL REFERENCES users (id), + data JSONB NOT NULL, + expires_at TIMESTAMPTZ NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +SELECT + diesel_manage_updated_at ('auth_sessions'); diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs index 26de85c..400bfe8 100644 --- a/server-new/src/db/mod.rs +++ b/server-new/src/db/mod.rs @@ -1,17 +1,18 @@ -use anyhow::Context; - -use crate::error::AppResult; - pub mod models; mod repositories; mod schema; /// Type of the database pool pub type DbPool = diesel_async::pooled_connection::deadpool::Pool; +/// Error when attempting to retrieve a connection from the pool +pub type DbPoolError = diesel_async::pooled_connection::deadpool::PoolError; /// Type of the database connection retrieved from the pool pub type DbConnection = diesel_async::pooled_connection::deadpool::Object; +/// Date/time format used in all database tables +pub type UtcDateTime = chrono::DateTime; + /// Wrapper around a database connection that gives access to the repositories, /// e.g. `UserRepository`, `ChatRepository`, etc. pub struct DbService { @@ -23,12 +24,16 @@ impl DbService { Self { cxn } } - pub async fn from_pool(pool: &DbPool) -> AppResult { - let cxn = pool.get().await.context("error retrieving DB connection")?; + pub async fn from_pool(pool: &DbPool) -> Result { + let cxn = pool.get().await?; Ok(Self::new(cxn)) } pub fn users(&mut self) -> repositories::UserRepository<'_> { repositories::UserRepository::new(&mut self.cxn) } + + pub fn sessions(&mut self) -> repositories::SessionRepository<'_> { + repositories::SessionRepository::new(&mut self.cxn) + } } diff --git a/server-new/src/db/models.rs b/server-new/src/db/models.rs index 10430d1..c225868 100644 --- a/server-new/src/db/models.rs +++ b/server-new/src/db/models.rs @@ -6,6 +6,7 @@ use crate::db::schema; // mod provider; // mod secret; // mod tool; +mod session; mod user; // pub use api_key::*; @@ -14,4 +15,5 @@ mod user; // pub use provider::*; // pub use secret::*; // pub use tool::*; +pub use session::*; pub use user::*; diff --git a/server-new/src/db/models/session.rs b/server-new/src/db/models/session.rs new file mode 100644 index 0000000..bacf0bf --- /dev/null +++ b/server-new/src/db/models/session.rs @@ -0,0 +1,65 @@ +use std::collections::HashMap; + +use diesel::{ + deserialize::{FromSql, FromSqlRow}, + expression::AsExpression, + prelude::*, + serialize::ToSql, +}; +use serde::{Deserialize, Serialize}; +use uuid::Uuid; + +use crate::db::{UtcDateTime, models::ChatRsUser}; + +#[derive(Identifiable, Associations, Queryable, Selectable)] +#[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] +#[diesel(table_name = super::schema::auth_sessions)] +pub struct ChatRsAuthSession { + pub id: Uuid, + pub user_id: Uuid, + pub data: AuthSessionData, + pub expires_at: UtcDateTime, +} + +#[derive(Insertable)] +#[diesel(table_name = super::schema::auth_sessions)] +pub struct NewChatRsAuthSession<'r> { + pub id: &'r Uuid, + pub user_id: &'r Uuid, + pub data: AuthSessionData, + pub expires_at: UtcDateTime, +} + +#[derive(AsChangeset)] +#[diesel(table_name = super::schema::auth_sessions)] +pub struct UpdateChatRsAuthSession { + pub data: AuthSessionData, + pub expires_at: UtcDateTime, +} + +#[derive(Debug, Serialize, Deserialize, FromSqlRow, AsExpression)] +#[diesel(sql_type = diesel::sql_types::Jsonb)] +pub struct AuthSessionData(pub HashMap); + +impl FromSql for AuthSessionData { + fn from_sql(bytes: diesel::pg::PgValue<'_>) -> diesel::deserialize::Result { + let value = + >::from_sql( + bytes, + )?; + Ok(serde_json::from_value(value)?) + } +} + +impl ToSql for AuthSessionData { + fn to_sql<'b>( + &'b self, + out: &mut diesel::serialize::Output<'b, '_, diesel::pg::Pg>, + ) -> diesel::serialize::Result { + let value = serde_json::to_value(self)?; + >::to_sql( + &value, + &mut out.reborrow(), + ) + } +} diff --git a/server-new/src/db/models/user.rs b/server-new/src/db/models/user.rs index 4383f7b..d2bd479 100644 --- a/server-new/src/db/models/user.rs +++ b/server-new/src/db/models/user.rs @@ -1,27 +1,24 @@ use diesel::prelude::*; use serde::Serialize; -use time::OffsetDateTime; +use serde_with::skip_serializing_none; use uuid::Uuid; +use crate::db::UtcDateTime; + +#[skip_serializing_none] #[derive(Identifiable, Queryable, Selectable, Serialize)] #[diesel(table_name = super::schema::users)] pub struct ChatRsUser { pub id: Uuid, pub name: String, - #[serde(skip_serializing_if = "Option::is_none")] pub avatar_url: Option, - #[serde(skip_serializing_if = "Option::is_none")] pub github_id: Option, - #[serde(skip_serializing_if = "Option::is_none")] pub google_id: Option, - #[serde(skip_serializing_if = "Option::is_none")] pub discord_id: Option, - #[serde(skip_serializing_if = "Option::is_none")] pub oidc_id: Option, - #[serde(skip_serializing_if = "Option::is_none")] pub sso_username: Option, - pub created_at: OffsetDateTime, - pub updated_at: OffsetDateTime, + pub created_at: chrono::DateTime, + pub updated_at: UtcDateTime, } #[derive(Insertable, Default)] diff --git a/server-new/src/db/repositories/mod.rs b/server-new/src/db/repositories/mod.rs index 5ba8673..7e3f251 100644 --- a/server-new/src/db/repositories/mod.rs +++ b/server-new/src/db/repositories/mod.rs @@ -1,3 +1,5 @@ +mod session; mod user; +pub use session::SessionRepository; pub use user::UserRepository; diff --git a/server-new/src/db/repositories/session.rs b/server-new/src/db/repositories/session.rs new file mode 100644 index 0000000..7dcbf77 --- /dev/null +++ b/server-new/src/db/repositories/session.rs @@ -0,0 +1,76 @@ +use std::collections::HashMap; + +use diesel::prelude::*; +use diesel_async::RunQueryDsl; +use uuid::Uuid; + +use crate::db::{ + DbConnection, UtcDateTime, + models::{AuthSessionData, ChatRsAuthSession, NewChatRsAuthSession, UpdateChatRsAuthSession}, + schema::auth_sessions, +}; + +pub struct SessionRepository<'a> { + db: &'a mut DbConnection, +} + +impl<'a> SessionRepository<'a> { + pub fn new(db: &'a mut DbConnection) -> Self { + Self { db } + } + + pub async fn create( + &mut self, + session_id: &Uuid, + user_id: &Uuid, + data: &HashMap, + expires_at: UtcDateTime, + ) -> QueryResult { + diesel::insert_into(auth_sessions::table) + .values(NewChatRsAuthSession { + id: session_id, + user_id, + data: AuthSessionData(data.to_owned()), + expires_at, + }) + .returning(ChatRsAuthSession::as_returning()) + .get_result(self.db) + .await + } + + pub async fn update( + &mut self, + session_id: &Uuid, + data: &HashMap, + expires_at: UtcDateTime, + ) -> QueryResult { + diesel::update(auth_sessions::table) + .filter(auth_sessions::id.eq(session_id)) + .set(UpdateChatRsAuthSession { + data: AuthSessionData(data.to_owned()), + expires_at, + }) + .returning(ChatRsAuthSession::as_returning()) + .get_result(self.db) + .await + } + + pub async fn find_by_id( + &mut self, + session_id: &Uuid, + ) -> QueryResult> { + auth_sessions::table + .find(session_id) + .select(ChatRsAuthSession::as_select()) + .first(self.db) + .await + .optional() + } + + pub async fn delete_by_id(&mut self, session_id: &Uuid) -> QueryResult { + diesel::delete(auth_sessions::table.find(session_id)) + .returning(auth_sessions::id) + .get_result(self.db) + .await + } +} diff --git a/server-new/src/db/repositories/user.rs b/server-new/src/db/repositories/user.rs index 9e211e4..8968284 100644 --- a/server-new/src/db/repositories/user.rs +++ b/server-new/src/db/repositories/user.rs @@ -3,11 +3,7 @@ use diesel::result::Error; use diesel_async::RunQueryDsl; use uuid::Uuid; -use crate::db::{ - DbConnection, - models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, - schema::users, -}; +use crate::db::{DbConnection, models::ChatRsUser, schema::users}; pub struct UserRepository<'a> { db: &'a mut DbConnection, diff --git a/server-new/src/db/schema.rs b/server-new/src/db/schema.rs index 2485dab..aa7d90d 100644 --- a/server-new/src/db/schema.rs +++ b/server-new/src/db/schema.rs @@ -19,6 +19,17 @@ diesel::table! { } } +diesel::table! { + auth_sessions (id) { + id -> Uuid, + user_id -> Uuid, + data -> Jsonb, + expires_at -> Timestamptz, + created_at -> Timestamptz, + updated_at -> Timestamptz, + } +} + diesel::table! { use diesel::sql_types::*; use super::sql_types::ChatMessageRole; @@ -123,6 +134,7 @@ diesel::table! { } diesel::joinable!(app_api_keys -> users (user_id)); +diesel::joinable!(auth_sessions -> users (user_id)); diesel::joinable!(chat_messages -> chat_sessions (session_id)); diesel::joinable!(chat_sessions -> users (user_id)); diesel::joinable!(external_api_tools -> users (user_id)); @@ -135,6 +147,7 @@ diesel::joinable!(system_tools -> users (user_id)); diesel::allow_tables_to_appear_in_same_query!( app_api_keys, + auth_sessions, chat_messages, chat_sessions, external_api_tools, diff --git a/server-new/src/error.rs b/server-new/src/error.rs index b75e3f3..dfb23e3 100644 --- a/server-new/src/error.rs +++ b/server-new/src/error.rs @@ -5,6 +5,8 @@ use axum::{ }; use serde::Serialize; +use crate::db::DbPoolError; + /// Global result type that can be used for API route handlers pub type AppResult = Result; @@ -16,15 +18,25 @@ pub struct AppError { source: Option, } +// Convenient error conversions impl From for AppError { fn from(error: anyhow::Error) -> Self { Self::internal(error) } } - impl From for AppError { fn from(error: diesel::result::Error) -> Self { - Self::internal(error.into()) + Self::internal(anyhow::Error::from(error).context("database error")) + } +} +impl From for AppError { + fn from(error: DbPoolError) -> Self { + Self::internal(anyhow::Error::from(error).context("pool error")) + } +} +impl From for AppError { + fn from(error: tower_sessions::session::Error) -> Self { + Self::internal(anyhow::Error::from(error).context("session error")) } } diff --git a/server-new/src/extractors/session.rs b/server-new/src/extractors/session.rs index c55271e..e92c7c5 100644 --- a/server-new/src/extractors/session.rs +++ b/server-new/src/extractors/session.rs @@ -3,25 +3,17 @@ use std::{ str::FromStr, }; -use anyhow::Context; use axum::{ extract::{ConnectInfo, FromRequestParts, OptionalFromRequestParts}, http::header, }; use serde::{Deserialize, Serialize}; -use time::UtcDateTime; use tower_sessions::Session; use uuid::Uuid; -use crate::{error::AppError, state::AppState}; +use crate::{db::UtcDateTime, error::AppError, state::AppState}; -/// The field used to store the user session data -const USER_SESSION_FIELD: &str = "sess"; -/// The field used to store the user session metadata -const SESSION_META_FIELD: &str = "meta"; - -/// Active session data. Beware when changing or adding to this struct, as it can -/// invalidate existing sessions. +/// Active user session data. /// /// This can be used as an extractor in route handlers: /// - If used as `Option`, will be `Some` if there is an active session @@ -45,51 +37,32 @@ impl UserSession { pub fn new(user_id: Uuid) -> Self { Self { user_id } } - - pub async fn init( - session: &Session, - meta: &SessionMeta, - user_id: &Uuid, - ) -> Result<(), AppError> { - session - .insert(USER_SESSION_FIELD, UserSession::new(user_id.clone())) - .await - .context("failed to initialize session")?; - session.insert(SESSION_META_FIELD, meta).await.ok(); - - Ok(()) - } } -impl OptionalFromRequestParts for UserSession { +impl OptionalFromRequestParts for UserSession { type Rejection = AppError; async fn from_request_parts( parts: &mut axum::http::request::Parts, - state: &S, + state: &AppState, ) -> Result, Self::Rejection> { let session = Session::from_request_parts(parts, state) .await .map_err(|(_, msg)| AppError::internal(anyhow::anyhow!(msg)))?; - match session.get::(USER_SESSION_FIELD).await { - Ok(Some(user_session)) => Ok(Some(user_session)), - Ok(None) => Ok(None), - Err(err) => Err(AppError::internal( - anyhow::Error::from(err).context("error while retrieving session"), - )), - } + state.auth_service().extract_user_session(session).await } } -impl FromRequestParts for UserSession { +impl FromRequestParts for UserSession { type Rejection = AppError; async fn from_request_parts( parts: &mut axum::http::request::Parts, - state: &S, + state: &AppState, ) -> Result { - match >::from_request_parts(parts, state).await? { + match >::from_request_parts(parts, state).await? + { Some(user_session) => Ok(user_session), None => Err(AppError::unauthorized()), } @@ -121,12 +94,11 @@ impl FromRequestParts for SessionMeta { .get(header::USER_AGENT) .and_then(|h| h.to_str().ok()) .map(|ua| ua.to_owned()); - let start_time = UtcDateTime::now(); Ok(Self { - start_time, ip, user_agent, + start_time: chrono::Utc::now(), }) } } diff --git a/server-new/src/plugins/database.rs b/server-new/src/plugins/database.rs index 58c32ad..51e992d 100644 --- a/server-new/src/plugins/database.rs +++ b/server-new/src/plugins/database.rs @@ -1,8 +1,7 @@ use anyhow::Context; use axum_plugin::AdHocPlugin; -use diesel::Connection; use diesel_async::{ - AsyncPgConnection, + AsyncMigrationHarness, AsyncPgConnection, pooled_connection::{AsyncDieselConnectionManager, deadpool::Pool}, }; use diesel_migrations::{EmbeddedMigrations, MigrationHarness}; @@ -15,28 +14,24 @@ pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Database") .on_init(async |mut state| { let app_config = state.get::().context("missing config")?; - let database_url = app_config.database_url.clone(); - - tokio::task::spawn_blocking(move || { - let mut cxn = diesel::PgConnection::establish(&database_url) - .context("Failed to connect to database")?; - tracing::info!("Connected to database at '{database_url}'"); - match cxn.run_pending_migrations(MIGRATIONS) { - Ok(run_migrations) => { - for migration in run_migrations { - tracing::info!("Migration run: '{migration}'"); - } - Ok(()) - } - Err(err) => Err(anyhow::anyhow!(err.to_string()).context("Migration failed")), - } - }) - .await??; let manager = AsyncDieselConnectionManager::::new(&app_config.database_url); let pool: DbPool = Pool::builder(manager).build()?; + let cxn = pool.get().await.context("failed to connect to database")?; + match AsyncMigrationHarness::new(cxn).run_pending_migrations(MIGRATIONS) { + Ok(run_migrations) if run_migrations.is_empty() => { + tracing::info!("No migrations to run"); + } + Ok(run_migrations) => { + for migration in run_migrations { + tracing::info!("Migration run: '{migration}'"); + } + } + Err(err) => anyhow::bail!(format!("Migrations failed: {err}")), + }; + state.insert(pool); Ok(state) }) diff --git a/server-new/src/plugins/session.rs b/server-new/src/plugins/session.rs index 1fff624..f86e32c 100644 --- a/server-new/src/plugins/session.rs +++ b/server-new/src/plugins/session.rs @@ -5,7 +5,7 @@ use tower_sessions::{ cookie::{Key, SameSite, time::Duration}, }; -use crate::state::AppState; +use crate::{services::SessionDbStore, state::AppState}; pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Session").on_setup(|router, state: &AppState| { @@ -15,8 +15,7 @@ pub fn plugin() -> AdHocPlugin { bail!("cookie_key must be at least 32 bytes"); } - // TODO change to a persistent store! - let session_store = tower_sessions::MemoryStore::default(); + let session_store = SessionDbStore::new(state.db_pool.clone()); let session_layer = SessionManagerLayer::new(session_store) .with_name(state.config.cookie_name.clone()) .with_expiry(Expiry::OnInactivity(Duration::seconds( diff --git a/server-new/src/routes/auth.rs b/server-new/src/routes/auth.rs index 90acbd8..530710c 100644 --- a/server-new/src/routes/auth.rs +++ b/server-new/src/routes/auth.rs @@ -1,5 +1,4 @@ -use anyhow::Context; -use axum::{extract::State, response::IntoResponse}; +use axum::{Json, extract::State, http::StatusCode, response::IntoResponse}; use crate::{ error::{AppError, AppResult}, @@ -18,14 +17,18 @@ async fn login_handler( maybe_user: Option, session: tower_sessions::Session, meta: SessionMeta, + State(state): State, ) -> AppResult { if maybe_user.is_some() { return Err(AppError::bad_request("already logged in")); } // TODO login handling logic - let user_id = uuid::Uuid::new_v4(); - UserSession::init(&session, &meta, &user_id).await?; + let user_id = uuid::Uuid::parse_str("6976658f-8eef-4a76-ad37-46243f463726").unwrap(); + state + .auth_service() + .init_session(&session, &meta, &user_id) + .await?; Ok(format!("Logged in as {user_id}")) } @@ -34,20 +37,15 @@ async fn get_user_handler( UserSession { user_id }: UserSession, State(state): State, ) -> AppResult { - state.user_service().get_user(&user_id).await?; + let user = state.auth_service().get_user(&user_id).await?; - Ok(format!("Logged in as {user_id}")) + Ok(Json(user)) } async fn logout_handler( - maybe_user: Option, session: tower_sessions::Session, + State(state): State, ) -> AppResult { - match maybe_user { - Some(_) => { - session.delete().await.context("error logging out")?; - Ok("Logged out") - } - None => Ok("Already logged out"), - } + state.auth_service().logout_user(&session).await?; + Ok(StatusCode::NO_CONTENT) } diff --git a/server-new/src/services/auth/mod.rs b/server-new/src/services/auth/mod.rs new file mode 100644 index 0000000..6f078f4 --- /dev/null +++ b/server-new/src/services/auth/mod.rs @@ -0,0 +1,68 @@ +use std::fmt::Debug; + +use crate::{ + db::{DbPool, DbService, models::ChatRsUser}, + error::{AppError, AppResult}, + extractors::session::{SessionMeta, UserSession}, +}; +use tower_sessions::Session; +use uuid::Uuid; + +mod session_store; +pub use session_store::SessionDbStore; + +/// The field used to store the user session data +const USER_ID_FIELD: &str = "user_id"; +/// The field used to store the user session metadata +const META_FIELD: &str = "meta"; + +pub struct AuthService<'a> { + db: &'a DbPool, +} + +impl<'a> Debug for AuthService<'a> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("AuthService").finish() + } +} + +impl<'a> AuthService<'a> { + pub fn new(db: &'a DbPool) -> Self { + Self { db } + } + + /// Initialize a new logged-in session for the given user + pub async fn init_session( + &self, + session: &Session, + meta: &SessionMeta, + user_id: &Uuid, + ) -> AppResult<()> { + session.insert(USER_ID_FIELD, user_id).await?; + session.insert(META_FIELD, meta).await?; + + Ok(()) + } + + /// Extract the current user session data + pub async fn extract_user_session(&self, session: Session) -> AppResult> { + let user_id = session.get::(USER_ID_FIELD).await?; + Ok(user_id.map(UserSession::new)) + } + + /// Get the user from the database with the given ID, or return + /// an internal error if not found + pub async fn get_user(&self, id: &Uuid) -> AppResult { + let mut db = DbService::from_pool(&self.db).await?; + match db.users().find_by_id(id).await? { + None => Err(AppError::internal(anyhow::anyhow!("user not found"))), + Some(user) => Ok(user), + } + } + + /// Logout the user, deleting the current session + pub async fn logout_user(&self, session: &Session) -> AppResult<()> { + session.flush().await?; + Ok(()) + } +} diff --git a/server-new/src/services/auth/session_store.rs b/server-new/src/services/auth/session_store.rs new file mode 100644 index 0000000..f1742a7 --- /dev/null +++ b/server-new/src/services/auth/session_store.rs @@ -0,0 +1,117 @@ +use async_trait::async_trait; +use tower_sessions::{ + SessionStore, + cookie::time::OffsetDateTime, + session::{Id, Record}, + session_store::{Error, Result}, +}; +use uuid::Uuid; + +use crate::db::{DbPool, DbService, UtcDateTime}; + +#[derive(Clone)] +pub struct SessionDbStore { + db: DbPool, +} + +impl SessionDbStore { + pub fn new(db: DbPool) -> Self { + Self { db } + } + + fn get_session_uuid(&self, id: &Id) -> Uuid { + Uuid::from_bytes(id.0.to_be_bytes()) + } + + fn convert_expiry(&self, time: OffsetDateTime) -> Result { + UtcDateTime::from_timestamp_secs(time.unix_timestamp()) + .ok_or_else(|| Error::Backend(format!("Invalid expiry: {time}"))) + } + + async fn get_db(&self) -> Result { + DbService::from_pool(&self.db) + .await + .map_err(|err| Error::Backend(err.to_string())) + } +} + +impl std::fmt::Debug for SessionDbStore { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("SessionDbStore").finish() + } +} + +#[async_trait] +impl SessionStore for SessionDbStore { + /// Creates a new session in the store with the provided session record. + async fn create(&self, record: &mut Record) -> Result<()> { + let session_id = self.get_session_uuid(&record.id); + let user_id_val = record + .data + .get(super::USER_ID_FIELD) + .ok_or_else(|| Error::Encode("no user id field".to_owned()))?; + let user_id: Uuid = serde_json::from_value(user_id_val.clone()) + .map_err(|_| Error::Encode("invalid user id field".to_owned()))?; + let expires_at = self.convert_expiry(record.expiry_date)?; + + let mut db = self.get_db().await?; + db.sessions() + .create(&session_id, &user_id, &record.data, expires_at) + .await + .map_err(|err| Error::Backend(err.to_string()))?; + + Ok(()) + } + + /// Saves the provided session record to the store. + /// + /// This method is intended for updating the state of an existing session. + async fn save(&self, record: &Record) -> Result<()> { + let session_id = self.get_session_uuid(&record.id); + let mut db = self.get_db().await?; + let expires_at = self.convert_expiry(record.expiry_date)?; + + db.sessions() + .update(&session_id, &record.data, expires_at) + .await + .map_err(|err| Error::Backend(err.to_string()))?; + + Ok(()) + } + + /// Loads an existing session record from the store using the provided ID. + /// + /// If a session with the given ID exists, it is returned. If the session + /// does not exist or has been invalidated (e.g., expired), `None` is + /// returned. + async fn load(&self, session_id: &Id) -> Result> { + let session_id = self.get_session_uuid(&session_id); + let mut db = self.get_db().await?; + + match db.sessions().find_by_id(&session_id).await { + Ok(Some(session)) => Ok(Some(Record { + id: Id(i128::from_be_bytes(session.id.into_bytes())), + data: session.data.0, + expiry_date: OffsetDateTime::from_unix_timestamp(session.expires_at.timestamp()) + .map_err(|err| Error::Backend(format!("Invalid expiry: {err}")))?, + })), + Ok(None) => Ok(None), + Err(err) => Err(Error::Backend(err.to_string())), + } + } + + /// Deletes a session record from the store using the provided ID. + /// + /// If the session exists, it is removed from the store. + async fn delete(&self, session_id: &Id) -> Result<()> { + let session_id = self.get_session_uuid(session_id); + let mut db = self.get_db().await?; + + db.sessions() + .delete_by_id(&session_id) + .await + .map_err(|err| Error::Backend(err.to_string()))?; + + Ok(()) + } +} diff --git a/server-new/src/services/mod.rs b/server-new/src/services/mod.rs index cbbc74c..07f4c79 100644 --- a/server-new/src/services/mod.rs +++ b/server-new/src/services/mod.rs @@ -1,3 +1,3 @@ -mod user; +mod auth; -pub use user::UserService; +pub use auth::{AuthService, SessionDbStore}; diff --git a/server-new/src/services/user.rs b/server-new/src/services/user.rs deleted file mode 100644 index 118b380..0000000 --- a/server-new/src/services/user.rs +++ /dev/null @@ -1,23 +0,0 @@ -use uuid::Uuid; - -use crate::{ - db::{DbPool, DbService, models::ChatRsUser}, - error::AppResult, -}; - -pub struct UserService<'a> { - db: &'a DbPool, -} - -impl<'a> UserService<'a> { - pub fn new(db: &'a DbPool) -> Self { - Self { db } - } - - pub async fn get_user(&self, id: &Uuid) -> AppResult> { - let mut db = DbService::from_pool(&self.db).await?; - let user = db.users().find_by_id(id).await?; - - Ok(user) - } -} diff --git a/server-new/src/state.rs b/server-new/src/state.rs index 1e87b99..f62f7cd 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -4,7 +4,7 @@ use std::{ops::Deref, sync::Arc}; use axum_plugin::{AppState, TypeMap}; -use crate::{config::AppConfig, db::DbPool, services::UserService}; +use crate::{config::AppConfig, db::DbPool, services::AuthService}; /// App state stored in the Axum router #[derive(Clone)] @@ -17,8 +17,8 @@ pub struct AppStateInner { } impl AppState { - pub fn user_service(&self) -> UserService<'_> { - UserService::new(&self.db_pool) + pub fn auth_service(&self) -> AuthService<'_> { + AuthService::new(&self.db_pool) } } From 41d0935526a4b4b7eec14763a6dc1234e8e813db Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 00:12:23 -0400 Subject: [PATCH 024/111] auth tweaks --- .../2026-06-22-224357-0000_add_auth_sessions/up.sql | 6 +++++- server-new/src/db/models/session.rs | 4 ++-- server-new/src/db/repositories/session.rs | 11 +++++++---- server-new/src/db/schema.rs | 2 +- server-new/src/extractors/session.rs | 3 ++- server-new/src/routes/auth.rs | 1 - server-new/src/services/auth/mod.rs | 11 +++++------ server-new/src/services/auth/session_store.rs | 10 +++++----- 8 files changed, 27 insertions(+), 21 deletions(-) diff --git a/server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/up.sql b/server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/up.sql index 3c7c368..f89d8f2 100644 --- a/server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/up.sql +++ b/server-new/migrations/2026-06-22-224357-0000_add_auth_sessions/up.sql @@ -1,6 +1,6 @@ CREATE TABLE auth_sessions ( id UUID PRIMARY KEY, - user_id UUID NOT NULL REFERENCES users (id), + user_id UUID NULL REFERENCES users (id), data JSONB NOT NULL, expires_at TIMESTAMPTZ NOT NULL, created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, @@ -9,3 +9,7 @@ CREATE TABLE auth_sessions ( SELECT diesel_manage_updated_at ('auth_sessions'); + +CREATE INDEX auth_sessions_user_id_idx ON auth_sessions (user_id); + +CREATE INDEX auth_sessions_expires_at_idx ON auth_sessions (expires_at); diff --git a/server-new/src/db/models/session.rs b/server-new/src/db/models/session.rs index bacf0bf..4972e41 100644 --- a/server-new/src/db/models/session.rs +++ b/server-new/src/db/models/session.rs @@ -16,7 +16,7 @@ use crate::db::{UtcDateTime, models::ChatRsUser}; #[diesel(table_name = super::schema::auth_sessions)] pub struct ChatRsAuthSession { pub id: Uuid, - pub user_id: Uuid, + pub user_id: Option, pub data: AuthSessionData, pub expires_at: UtcDateTime, } @@ -25,7 +25,7 @@ pub struct ChatRsAuthSession { #[diesel(table_name = super::schema::auth_sessions)] pub struct NewChatRsAuthSession<'r> { pub id: &'r Uuid, - pub user_id: &'r Uuid, + pub user_id: Option<&'r Uuid>, pub data: AuthSessionData, pub expires_at: UtcDateTime, } diff --git a/server-new/src/db/repositories/session.rs b/server-new/src/db/repositories/session.rs index 7dcbf77..ecfad57 100644 --- a/server-new/src/db/repositories/session.rs +++ b/server-new/src/db/repositories/session.rs @@ -22,7 +22,7 @@ impl<'a> SessionRepository<'a> { pub async fn create( &mut self, session_id: &Uuid, - user_id: &Uuid, + user_id: Option<&Uuid>, data: &HashMap, expires_at: UtcDateTime, ) -> QueryResult { @@ -55,22 +55,25 @@ impl<'a> SessionRepository<'a> { .await } - pub async fn find_by_id( + /// Find an active (not expired) session by ID + pub async fn find_active_by_id( &mut self, session_id: &Uuid, ) -> QueryResult> { auth_sessions::table .find(session_id) + .filter(auth_sessions::expires_at.gt(diesel::dsl::now)) .select(ChatRsAuthSession::as_select()) .first(self.db) .await .optional() } - pub async fn delete_by_id(&mut self, session_id: &Uuid) -> QueryResult { + /// Delete a session by ID. Won't return an error if it does not exist. + pub async fn delete_by_id(&mut self, session_id: &Uuid) -> QueryResult { diesel::delete(auth_sessions::table.find(session_id)) .returning(auth_sessions::id) - .get_result(self.db) + .execute(self.db) .await } } diff --git a/server-new/src/db/schema.rs b/server-new/src/db/schema.rs index aa7d90d..7b7cee5 100644 --- a/server-new/src/db/schema.rs +++ b/server-new/src/db/schema.rs @@ -22,7 +22,7 @@ diesel::table! { diesel::table! { auth_sessions (id) { id -> Uuid, - user_id -> Uuid, + user_id -> Nullable, data -> Jsonb, expires_at -> Timestamptz, created_at -> Timestamptz, diff --git a/server-new/src/extractors/session.rs b/server-new/src/extractors/session.rs index e92c7c5..d1157d5 100644 --- a/server-new/src/extractors/session.rs +++ b/server-new/src/extractors/session.rs @@ -49,8 +49,9 @@ impl OptionalFromRequestParts for UserSession { let session = Session::from_request_parts(parts, state) .await .map_err(|(_, msg)| AppError::internal(anyhow::anyhow!(msg)))?; + let user_id = state.auth_service().extract_user_id(session).await?; - state.auth_service().extract_user_session(session).await + Ok(user_id.map(UserSession::new)) } } diff --git a/server-new/src/routes/auth.rs b/server-new/src/routes/auth.rs index 530710c..f3367dd 100644 --- a/server-new/src/routes/auth.rs +++ b/server-new/src/routes/auth.rs @@ -38,7 +38,6 @@ async fn get_user_handler( State(state): State, ) -> AppResult { let user = state.auth_service().get_user(&user_id).await?; - Ok(Json(user)) } diff --git a/server-new/src/services/auth/mod.rs b/server-new/src/services/auth/mod.rs index 6f078f4..0fc4eaa 100644 --- a/server-new/src/services/auth/mod.rs +++ b/server-new/src/services/auth/mod.rs @@ -3,7 +3,7 @@ use std::fmt::Debug; use crate::{ db::{DbPool, DbService, models::ChatRsUser}, error::{AppError, AppResult}, - extractors::session::{SessionMeta, UserSession}, + extractors::session::SessionMeta, }; use tower_sessions::Session; use uuid::Uuid; @@ -11,7 +11,7 @@ use uuid::Uuid; mod session_store; pub use session_store::SessionDbStore; -/// The field used to store the user session data +/// The field used to store the user ID in the session const USER_ID_FIELD: &str = "user_id"; /// The field used to store the user session metadata const META_FIELD: &str = "meta"; @@ -44,10 +44,9 @@ impl<'a> AuthService<'a> { Ok(()) } - /// Extract the current user session data - pub async fn extract_user_session(&self, session: Session) -> AppResult> { - let user_id = session.get::(USER_ID_FIELD).await?; - Ok(user_id.map(UserSession::new)) + /// Extract the current user ID if this is an active user session + pub async fn extract_user_id(&self, session: Session) -> AppResult> { + Ok(session.get::(USER_ID_FIELD).await?) } /// Get the user from the database with the given ID, or return diff --git a/server-new/src/services/auth/session_store.rs b/server-new/src/services/auth/session_store.rs index f1742a7..caa4daa 100644 --- a/server-new/src/services/auth/session_store.rs +++ b/server-new/src/services/auth/session_store.rs @@ -46,17 +46,17 @@ impl SessionStore for SessionDbStore { /// Creates a new session in the store with the provided session record. async fn create(&self, record: &mut Record) -> Result<()> { let session_id = self.get_session_uuid(&record.id); - let user_id_val = record + let user_id = record .data .get(super::USER_ID_FIELD) - .ok_or_else(|| Error::Encode("no user id field".to_owned()))?; - let user_id: Uuid = serde_json::from_value(user_id_val.clone()) + .map(|val| serde_json::from_value::(val.clone())) + .transpose() .map_err(|_| Error::Encode("invalid user id field".to_owned()))?; let expires_at = self.convert_expiry(record.expiry_date)?; let mut db = self.get_db().await?; db.sessions() - .create(&session_id, &user_id, &record.data, expires_at) + .create(&session_id, user_id.as_ref(), &record.data, expires_at) .await .map_err(|err| Error::Backend(err.to_string()))?; @@ -88,7 +88,7 @@ impl SessionStore for SessionDbStore { let session_id = self.get_session_uuid(&session_id); let mut db = self.get_db().await?; - match db.sessions().find_by_id(&session_id).await { + match db.sessions().find_active_by_id(&session_id).await { Ok(Some(session)) => Ok(Some(Record { id: Id(i128::from_be_bytes(session.id.into_bytes())), data: session.data.0, From be2d7222819c6893ec21df50cec4d020fc2220f8 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 00:23:02 -0400 Subject: [PATCH 025/111] add diesel jsonb derive macro --- server-new/Cargo.lock | 10 +++++ server-new/Cargo.toml | 1 + server-new/diesel-jsonb-derive/Cargo.toml | 14 +++++++ server-new/diesel-jsonb-derive/src/lib.rs | 46 +++++++++++++++++++++++ server-new/src/db/models/session.rs | 33 ++-------------- 5 files changed, 74 insertions(+), 30 deletions(-) create mode 100644 server-new/diesel-jsonb-derive/Cargo.toml create mode 100644 server-new/diesel-jsonb-derive/src/lib.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index bf283cc..c09d083 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -458,6 +458,15 @@ dependencies = [ "tokio-postgres", ] +[[package]] +name = "diesel-jsonb-derive" +version = "0.1.0" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "diesel_derives" version = "2.3.9" @@ -1352,6 +1361,7 @@ dependencies = [ "chrono", "diesel", "diesel-async", + "diesel-jsonb-derive", "diesel_migrations", "dotenvy", "figment", diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 72e75a6..f550224 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -23,6 +23,7 @@ diesel = { default-features = false, features = ["chrono", "serde_json", "uuid"] } +diesel-jsonb-derive = { path = "./diesel-jsonb-derive" } diesel-async = { version = "0.9.2", features = ["deadpool", "migrations", "postgres"] diff --git a/server-new/diesel-jsonb-derive/Cargo.toml b/server-new/diesel-jsonb-derive/Cargo.toml new file mode 100644 index 0000000..fc9db4e --- /dev/null +++ b/server-new/diesel-jsonb-derive/Cargo.toml @@ -0,0 +1,14 @@ +[package] +name = "diesel-jsonb-derive" +version = "0.1.0" +edition = "2024" +description = "Internal derive macro for Diesel JSONB serde conversion" +publish = false + +[lib] +proc-macro = true + +[dependencies] +proc-macro2 = "1" +quote = "1" +syn = { version = "2", features = ["derive"] } diff --git a/server-new/diesel-jsonb-derive/src/lib.rs b/server-new/diesel-jsonb-derive/src/lib.rs new file mode 100644 index 0000000..2b2db5f --- /dev/null +++ b/server-new/diesel-jsonb-derive/src/lib.rs @@ -0,0 +1,46 @@ +use proc_macro::TokenStream; +use quote::quote; +use syn::{DeriveInput, parse_macro_input}; + +#[proc_macro_derive(AsJsonb)] +pub fn derive_as_jsonb(input: TokenStream) -> TokenStream { + let input = parse_macro_input!(input as DeriveInput); + let ident = input.ident; + let generics = input.generics; + let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); + + quote! { + impl #impl_generics diesel::deserialize::FromSql + for #ident #ty_generics + #where_clause + { + fn from_sql(bytes: diesel::pg::PgValue<'_>) -> diesel::deserialize::Result { + let value = + >::from_sql(bytes)?; + + Ok(serde_json::from_value(value)?) + } + } + + impl #impl_generics diesel::serialize::ToSql + for #ident #ty_generics + #where_clause + { + fn to_sql<'b>( + &'b self, + out: &mut diesel::serialize::Output<'b, '_, diesel::pg::Pg>, + ) -> diesel::serialize::Result { + let value = serde_json::to_value(self)?; + + >::to_sql(&value, &mut out.reborrow()) + } + } + } + .into() +} diff --git a/server-new/src/db/models/session.rs b/server-new/src/db/models/session.rs index 4972e41..3084efa 100644 --- a/server-new/src/db/models/session.rs +++ b/server-new/src/db/models/session.rs @@ -1,11 +1,7 @@ use std::collections::HashMap; -use diesel::{ - deserialize::{FromSql, FromSqlRow}, - expression::AsExpression, - prelude::*, - serialize::ToSql, -}; +use diesel::{deserialize::FromSqlRow, expression::AsExpression, prelude::*}; +use diesel_jsonb_derive::AsJsonb; use serde::{Deserialize, Serialize}; use uuid::Uuid; @@ -37,29 +33,6 @@ pub struct UpdateChatRsAuthSession { pub expires_at: UtcDateTime, } -#[derive(Debug, Serialize, Deserialize, FromSqlRow, AsExpression)] +#[derive(Debug, Serialize, Deserialize, FromSqlRow, AsExpression, AsJsonb)] #[diesel(sql_type = diesel::sql_types::Jsonb)] pub struct AuthSessionData(pub HashMap); - -impl FromSql for AuthSessionData { - fn from_sql(bytes: diesel::pg::PgValue<'_>) -> diesel::deserialize::Result { - let value = - >::from_sql( - bytes, - )?; - Ok(serde_json::from_value(value)?) - } -} - -impl ToSql for AuthSessionData { - fn to_sql<'b>( - &'b self, - out: &mut diesel::serialize::Output<'b, '_, diesel::pg::Pg>, - ) -> diesel::serialize::Result { - let value = serde_json::to_value(self)?; - >::to_sql( - &value, - &mut out.reborrow(), - ) - } -} From 64c2502902e4e16f6946b6a4150991497ad859ca Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 00:30:12 -0400 Subject: [PATCH 026/111] re-arrange --- server-new/Cargo.toml | 2 +- server-new/Dockerfile | 4 ++-- server-new/{ => crates}/diesel-jsonb-derive/Cargo.toml | 0 server-new/{ => crates}/diesel-jsonb-derive/src/lib.rs | 0 4 files changed, 3 insertions(+), 3 deletions(-) rename server-new/{ => crates}/diesel-jsonb-derive/Cargo.toml (100%) rename server-new/{ => crates}/diesel-jsonb-derive/src/lib.rs (100%) diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index f550224..5fb42a6 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -23,11 +23,11 @@ diesel = { default-features = false, features = ["chrono", "serde_json", "uuid"] } -diesel-jsonb-derive = { path = "./diesel-jsonb-derive" } diesel-async = { version = "0.9.2", features = ["deadpool", "migrations", "postgres"] } +diesel-jsonb-derive = { path = "./crates/diesel-jsonb-derive" } diesel_migrations = { version = "2.3.2", features = ["postgres"] } dotenvy = "0.15.7" figment = { version = "0.10.19", features = ["env"] } diff --git a/server-new/Dockerfile b/server-new/Dockerfile index 88b7054..1ebfd5b 100644 --- a/server-new/Dockerfile +++ b/server-new/Dockerfile @@ -8,9 +8,9 @@ WORKDIR /app # Copy all necessary files to build the server COPY Cargo.lock Cargo.toml ./ +COPY ./crates ./crates +COPY ./migrations ./migrations COPY ./src ./src -# COPY ./migrations ./migrations -# etc... ARG pkg=rs-chat-api diff --git a/server-new/diesel-jsonb-derive/Cargo.toml b/server-new/crates/diesel-jsonb-derive/Cargo.toml similarity index 100% rename from server-new/diesel-jsonb-derive/Cargo.toml rename to server-new/crates/diesel-jsonb-derive/Cargo.toml diff --git a/server-new/diesel-jsonb-derive/src/lib.rs b/server-new/crates/diesel-jsonb-derive/src/lib.rs similarity index 100% rename from server-new/diesel-jsonb-derive/src/lib.rs rename to server-new/crates/diesel-jsonb-derive/src/lib.rs From b47750a3d324437153f29a9ccef2b96beb175b86 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 01:22:13 -0400 Subject: [PATCH 027/111] refactor config, add redis --- server-new/.env.example | 12 +- server-new/Cargo.lock | 626 ++++++++++++++++++++++++++- server-new/Cargo.toml | 11 +- server-new/Dockerfile | 3 +- server-new/config.toml | 19 + server-new/src/config.rs | 72 ++- server-new/src/extractors/session.rs | 1 + server-new/src/lib.rs | 1 + server-new/src/main.rs | 8 +- server-new/src/plugins/database.rs | 3 +- server-new/src/plugins/logging.rs | 2 +- server-new/src/plugins/mod.rs | 1 + server-new/src/plugins/redis.rs | 33 ++ server-new/src/plugins/security.rs | 4 +- server-new/src/plugins/session.rs | 6 +- server-new/src/state.rs | 1 + 16 files changed, 728 insertions(+), 75 deletions(-) create mode 100644 server-new/config.toml create mode 100644 server-new/src/plugins/redis.rs diff --git a/server-new/.env.example b/server-new/.env.example index 8c718da..a863a03 100644 --- a/server-new/.env.example +++ b/server-new/.env.example @@ -1,14 +1,12 @@ # Server Configuration -RS_CHAT_HOST=127.0.0.1 -RS_CHAT_PORT=8080 +RS_CHAT_SERVER__PORT=8080 -# Database -RS_CHAT_DATABASE_URL=postgres://postgres:postgres@localhost/postgres +# Database (2nd one is needed for Diesel CLI) +RS_CHAT_DATABASE__URL=postgres://postgres:postgres@localhost/postgres DATABASE_URL=postgres://postgres:postgres@localhost/postgres # Auth -RS_CHAT_COOKIE_KEY= # hex secret >=32 bytes, e.g. `openssl rand --hex 32` +RS_CHAT_AUTH__COOKIE_KEY= # hex secret >=32 bytes, e.g. `openssl rand --hex 32` # Logging -RS_CHAT_LOG_LEVEL=info -RS_CHAT_REQUEST_ID_HEADER=x-request-id +RS_CHAT_SERVER__LOG_LEVEL=info diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index c09d083..edca47d 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -61,6 +61,15 @@ version = "1.0.102" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" +[[package]] +name = "arc-swap" +version = "1.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6a3a1fd6f75306b68087b831f025c712524bcb19aad54e557b1129cfa0a2b207" +dependencies = [ + "rustversion", +] + [[package]] name = "async-trait" version = "0.1.89" @@ -225,6 +234,16 @@ version = "1.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8ae3f5d315924270530207e2a68396c3cc547f6dca3fbdca317cfb1a51edb593" +[[package]] +name = "bytes-utils" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7dafe3a8757b027e2be6e4e5601ed563c55989fcf1546e933c66c8eb3a058d35" +dependencies = [ + "bytes", + "either", +] + [[package]] name = "cc" version = "1.2.65" @@ -281,6 +300,12 @@ dependencies = [ "version_check", ] +[[package]] +name = "cookie-factory" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "396de984970346b0d9e93d1415082923c679e5ae5c3ee3dcbd104f5610af126b" + [[package]] name = "core-foundation-sys" version = "0.8.7" @@ -296,6 +321,12 @@ dependencies = [ "libc", ] +[[package]] +name = "crc16" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "338089f42c427b86394a5ee60ff321da23a5c89c9d89514c829687b26359fcff" + [[package]] name = "crossbeam-channel" version = "0.5.15" @@ -511,6 +542,17 @@ dependencies = [ "subtle", ] +[[package]] +name = "displaydoc" +version = "0.2.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ac70aa55017e108007fbaf5aa0f54b021c98f92ff8af59d42eda9da96e3dd4f" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "dotenvy" version = "0.15.7" @@ -543,6 +585,12 @@ version = "1.16.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + [[package]] name = "errno" version = "0.3.14" @@ -550,7 +598,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys", + "windows-sys 0.61.2", ] [[package]] @@ -568,6 +616,7 @@ dependencies = [ "atomic", "pear", "serde", + "toml 0.8.23", "uncased", "version_check", ] @@ -578,6 +627,15 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" +[[package]] +name = "float-cmp" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b09cf3155332e944990140d967ff5eceb70df778b34f77d8075db46e4704e6d8" +dependencies = [ + "num-traits", +] + [[package]] name = "fnv" version = "1.0.7" @@ -593,6 +651,43 @@ dependencies = [ "percent-encoding", ] +[[package]] +name = "fred" +version = "10.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3a7b2fd0f08b23315c13b6156f971aeedb6f75fb16a29ac1872d2eabccc1490e" +dependencies = [ + "arc-swap", + "async-trait", + "bytes", + "bytes-utils", + "float-cmp", + "fred-macros", + "futures", + "log", + "parking_lot", + "rand 0.8.6", + "redis-protocol", + "semver", + "socket2 0.5.10", + "tokio", + "tokio-stream", + "tokio-util", + "url", + "urlencoding", +] + +[[package]] +name = "fred-macros" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1458c6e22d36d61507034d5afecc64f105c1d39712b7ac6ec3b352c423f715cc" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "futures" version = "0.3.32" @@ -735,6 +830,12 @@ dependencies = [ "polyval", ] +[[package]] +name = "hashbrown" +version = "0.17.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" + [[package]] name = "heck" version = "0.5.0" @@ -881,12 +982,125 @@ dependencies = [ "cc", ] +[[package]] +name = "icu_collections" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2984d1cd16c883d7935b9e07e44071dca8d917fd52ecc02c04d5fa0b5a3f191c" +dependencies = [ + "displaydoc", + "potential_utf", + "utf8_iter", + "yoke", + "zerofrom", + "zerovec", +] + +[[package]] +name = "icu_locale_core" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92219b62b3e2b4d88ac5119f8904c10f8f61bf7e95b640d25ba3075e6cac2c29" +dependencies = [ + "displaydoc", + "litemap", + "tinystr", + "writeable", + "zerovec", +] + +[[package]] +name = "icu_normalizer" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c56e5ee99d6e3d33bd91c5d85458b6005a22140021cc324cea84dd0e72cff3b4" +dependencies = [ + "icu_collections", + "icu_normalizer_data", + "icu_properties", + "icu_provider", + "smallvec", + "zerovec", +] + +[[package]] +name = "icu_normalizer_data" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da3be0ae77ea334f4da67c12f149704f19f81d1adf7c51cf482943e84a2bad38" + +[[package]] +name = "icu_properties" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bee3b67d0ea5c2cca5003417989af8996f8604e34fb9ddf96208a033901e70de" +dependencies = [ + "icu_collections", + "icu_locale_core", + "icu_properties_data", + "icu_provider", + "zerotrie", + "zerovec", +] + +[[package]] +name = "icu_properties_data" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e2bbb201e0c04f7b4b3e14382af113e17ba4f63e2c9d2ee626b720cbce54a14" + +[[package]] +name = "icu_provider" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "139c4cf31c8b5f33d7e199446eff9c1e02decfc2f0eec2c8d71f65befa45b421" +dependencies = [ + "displaydoc", + "icu_locale_core", + "writeable", + "yoke", + "zerofrom", + "zerotrie", + "zerovec", +] + [[package]] name = "ident_case" version = "1.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b9e0384b61958566e926dc50660321d12159025e767c18e043daf26b70104c39" +[[package]] +name = "idna" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3b0875f23caa03898994f6ddc501886a45c7d3d62d04d2d90788d47be1b1e4de" +dependencies = [ + "idna_adapter", + "smallvec", + "utf8_iter", +] + +[[package]] +name = "idna_adapter" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb68373c0d6620ef8105e855e7745e18b0d00d3bdb07fb532e434244cdb9a714" +dependencies = [ + "icu_normalizer", + "icu_properties", +] + +[[package]] +name = "indexmap" +version = "2.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" +dependencies = [ + "equivalent", + "hashbrown", +] + [[package]] name = "inlinable_string" version = "0.1.15" @@ -940,6 +1154,12 @@ dependencies = [ "libc", ] +[[package]] +name = "litemap" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92daf443525c4cce67b150400bc2316076100ce0b3686209eb8cf3c31612e6f0" + [[package]] name = "lock_api" version = "0.4.14" @@ -947,6 +1167,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "224399e74b87b5f3557511d98dff8b14089b3dadafcab6bb93eab67d3aace965" dependencies = [ "scopeguard", + "serde", ] [[package]] @@ -993,7 +1214,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "36c791ecdf977c99f45f23280405d7723727470f6689a5e6dbf513ac547ae10d" dependencies = [ "serde", - "toml", + "toml 0.9.12+spec-1.1.0", ] [[package]] @@ -1013,6 +1234,12 @@ version = "0.3.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a" +[[package]] +name = "minimal-lexical" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" + [[package]] name = "mio" version = "1.2.1" @@ -1021,7 +1248,17 @@ checksum = "02bd0af71c67b473010cbbc60715ee815645a4dc942899111f494b4b737d6fda" dependencies = [ "libc", "wasi 0.11.1+wasi-snapshot-preview1", - "windows-sys", + "windows-sys 0.61.2", +] + +[[package]] +name = "nom" +version = "7.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d273983c5a657a70a3e8f2a01329822f3b8c8172b73826411a55751e404a0a4a" +dependencies = [ + "memchr", + "minimal-lexical", ] [[package]] @@ -1030,7 +1267,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys", + "windows-sys 0.61.2", ] [[package]] @@ -1206,6 +1443,15 @@ dependencies = [ "postgres-protocol", ] +[[package]] +name = "potential_utf" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0103b1cef7ec0cf76490e969665504990193874ea05c85ff9bab8b911d0a0564" +dependencies = [ + "zerovec", +] + [[package]] name = "powerfmt" version = "0.2.0" @@ -1323,6 +1569,20 @@ dependencies = [ "getrandom 0.3.4", ] +[[package]] +name = "redis-protocol" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9cdba59219406899220fc4cdfd17a95191ba9c9afb719b5fa5a083d63109a9f1" +dependencies = [ + "bytes", + "bytes-utils", + "cookie-factory", + "crc16", + "log", + "nom", +] + [[package]] name = "redox_syscall" version = "0.5.18" @@ -1349,6 +1609,25 @@ version = "0.8.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" +[[package]] +name = "rmp" +version = "0.8.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4ba8be72d372b2c9b35542551678538b562e7cf86c3315773cae48dfbfe7790c" +dependencies = [ + "num-traits", +] + +[[package]] +name = "rmp-serde" +version = "1.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72f81bee8c8ef9b577d1681a70ebbc962c232461e397b22c208c43c04b67a155" +dependencies = [ + "rmp", + "serde", +] + [[package]] name = "rs-chat-api" version = "0.1.0" @@ -1365,6 +1644,7 @@ dependencies = [ "diesel_migrations", "dotenvy", "figment", + "fred", "hex", "serde", "serde_json", @@ -1373,6 +1653,7 @@ dependencies = [ "tower", "tower-http", "tower-sessions", + "tower-sessions-redis-store", "tracing", "tracing-appender", "tracing-subscriber", @@ -1403,6 +1684,12 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" +[[package]] +name = "semver" +version = "1.0.28" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a7852d02fc848982e0c167ef163aaff9cd91dc640ba85e263cb1ce46fae51cd" + [[package]] name = "serde" version = "1.0.228" @@ -1457,6 +1744,15 @@ dependencies = [ "serde_core", ] +[[package]] +name = "serde_spanned" +version = "0.6.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bf41e0cfaf7226dca15e8197172c295a782857fcb97fad1808a166870dee75a3" +dependencies = [ + "serde", +] + [[package]] name = "serde_spanned" version = "1.1.1" @@ -1554,6 +1850,16 @@ version = "1.15.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8ed6a63f02c8539c91a8685a86f4099661ba3da017932f6ebbea6de3f0fa7c90" +[[package]] +name = "socket2" +version = "0.5.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e22376abed350d73dd1cd119b57ffccad95b4e585a7cda43e286245ce23c0678" +dependencies = [ + "libc", + "windows-sys 0.52.0", +] + [[package]] name = "socket2" version = "0.6.4" @@ -1561,9 +1867,15 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "52d1cfed4120b4d927bf7c0f86d2087a4a7d6027c906d9f9d525a80573b9be51" dependencies = [ "libc", - "windows-sys", + "windows-sys 0.61.2", ] +[[package]] +name = "stable_deref_trait" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" + [[package]] name = "stringprep" version = "0.1.5" @@ -1610,6 +1922,17 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" +[[package]] +name = "synstructure" +version = "0.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "thiserror" version = "2.0.18" @@ -1669,6 +1992,16 @@ dependencies = [ "time-core", ] +[[package]] +name = "tinystr" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8323304221c2a851516f22236c5722a72eaa19749016521d6dff0824447d96d" +dependencies = [ + "displaydoc", + "zerovec", +] + [[package]] name = "tinyvec" version = "1.11.0" @@ -1695,9 +2028,9 @@ dependencies = [ "mio", "pin-project-lite", "signal-hook-registry", - "socket2", + "socket2 0.6.4", "tokio-macros", - "windows-sys", + "windows-sys 0.61.2", ] [[package]] @@ -1731,12 +2064,23 @@ dependencies = [ "postgres-protocol", "postgres-types", "rand 0.9.4", - "socket2", + "socket2 0.6.4", "tokio", "tokio-util", "whoami", ] +[[package]] +name = "tokio-stream" +version = "0.1.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32da49809aab5c3bc678af03902d4ccddea2a87d028d86392a4b1560c6906c70" +dependencies = [ + "futures-core", + "pin-project-lite", + "tokio", +] + [[package]] name = "tokio-util" version = "0.7.18" @@ -1750,6 +2094,18 @@ dependencies = [ "tokio", ] +[[package]] +name = "toml" +version = "0.8.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc1beb996b9d83529a9e75c17a1686767d148d70663143c7854d8b4a09ced362" +dependencies = [ + "serde", + "serde_spanned 0.6.9", + "toml_datetime 0.6.11", + "toml_edit", +] + [[package]] name = "toml" version = "0.9.12+spec-1.1.0" @@ -1757,12 +2113,21 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf92845e79fc2e2def6a5d828f0801e29a2f8acc037becc5ab08595c7d5e9863" dependencies = [ "serde_core", - "serde_spanned", - "toml_datetime", + "serde_spanned 1.1.1", + "toml_datetime 0.7.5+spec-1.1.0", "toml_parser", "winnow 0.7.15", ] +[[package]] +name = "toml_datetime" +version = "0.6.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22cddaf88f4fbc13c51aebbf5f8eceb5c7c5a9da2ac40a13519eb5b0a0e8f11c" +dependencies = [ + "serde", +] + [[package]] name = "toml_datetime" version = "0.7.5+spec-1.1.0" @@ -1772,6 +2137,20 @@ dependencies = [ "serde_core", ] +[[package]] +name = "toml_edit" +version = "0.22.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41fe8c660ae4257887cf66394862d21dbca4a6ddd26f04a3560410406a2f819a" +dependencies = [ + "indexmap", + "serde", + "serde_spanned 0.6.9", + "toml_datetime 0.6.11", + "toml_write", + "winnow 0.7.15", +] + [[package]] name = "toml_parser" version = "1.1.2+spec-1.1.0" @@ -1781,6 +2160,12 @@ dependencies = [ "winnow 1.0.3", ] +[[package]] +name = "toml_write" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5d99f8c9a7727884afe522e9bd5edbfc91a3312b36a77b5fb8926e4c31a41801" + [[package]] name = "tower" version = "0.5.3" @@ -1858,11 +2243,31 @@ dependencies = [ "tower-cookies", "tower-layer", "tower-service", - "tower-sessions-core", + "tower-sessions-core 0.15.0", "tower-sessions-memory-store", "tracing", ] +[[package]] +name = "tower-sessions-core" +version = "0.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ce8cce604865576b7751b7a6bc3058f754569a60d689328bb74c52b1d87e355b" +dependencies = [ + "async-trait", + "base64", + "futures", + "http", + "parking_lot", + "rand 0.8.6", + "serde", + "serde_json", + "thiserror", + "time", + "tokio", + "tracing", +] + [[package]] name = "tower-sessions-core" version = "0.15.0" @@ -1893,7 +2298,21 @@ dependencies = [ "async-trait", "time", "tokio", - "tower-sessions-core", + "tower-sessions-core 0.15.0", +] + +[[package]] +name = "tower-sessions-redis-store" +version = "0.16.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e15b774f3d46625a27a8ac1238ecd73c8bd50013244e2de004026e161aad728" +dependencies = [ + "async-trait", + "fred", + "rmp-serde", + "thiserror", + "time", + "tower-sessions-core 0.14.0", ] [[package]] @@ -2045,6 +2464,30 @@ dependencies = [ "subtle", ] +[[package]] +name = "url" +version = "2.5.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff67a8a4397373c3ef660812acab3268222035010ab8680ec4215f38ba3d0eed" +dependencies = [ + "form_urlencoded", + "idna", + "percent-encoding", + "serde", +] + +[[package]] +name = "urlencoding" +version = "2.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "daf8dba3b7eb870caf1ddeed7bc9d2a049f3cfdfae7cb521b087cc33ae4c49da" + +[[package]] +name = "utf8_iter" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" + [[package]] name = "uuid" version = "1.23.3" @@ -2229,6 +2672,15 @@ dependencies = [ "windows-link", ] +[[package]] +name = "windows-sys" +version = "0.52.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" +dependencies = [ + "windows-targets", +] + [[package]] name = "windows-sys" version = "0.61.2" @@ -2238,11 +2690,78 @@ dependencies = [ "windows-link", ] +[[package]] +name = "windows-targets" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" +dependencies = [ + "windows_aarch64_gnullvm", + "windows_aarch64_msvc", + "windows_i686_gnu", + "windows_i686_gnullvm", + "windows_i686_msvc", + "windows_x86_64_gnu", + "windows_x86_64_gnullvm", + "windows_x86_64_msvc", +] + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" + +[[package]] +name = "windows_i686_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" + +[[package]] +name = "windows_i686_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" + [[package]] name = "winnow" version = "0.7.15" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945" +dependencies = [ + "memchr", +] [[package]] name = "winnow" @@ -2256,12 +2775,41 @@ version = "0.57.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" +[[package]] +name = "writeable" +version = "0.6.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4" + [[package]] name = "yansi" version = "1.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cfe53a6657fd280eaa890a3bc59152892ffa3e30101319d168b781ed6529b049" +[[package]] +name = "yoke" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "709fe23a0424b6a435d82152b1bd3fdfb0833487d5fa90d05d42762a9891fef5" +dependencies = [ + "stable_deref_trait", + "yoke-derive", + "zerofrom", +] + +[[package]] +name = "yoke-derive" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "de844c262c8848816172cef550288e7dc6c7b7814b4ee56b3e1553f275f1858e" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "synstructure", +] + [[package]] name = "zerocopy" version = "0.8.52" @@ -2282,6 +2830,60 @@ dependencies = [ "syn", ] +[[package]] +name = "zerofrom" +version = "0.1.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ec05a11813ea801ff6d75110ad09cd0824ddba17dfe17128ea0d5f68e6c5272" +dependencies = [ + "zerofrom-derive", +] + +[[package]] +name = "zerofrom-derive" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "11532158c46691caf0f2593ea8358fed6bbf68a0315e80aae9bd41fbade684a1" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "synstructure", +] + +[[package]] +name = "zerotrie" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0f9152d31db0792fa83f70fb2f83148effb5c1f5b8c7686c3459e361d9bc20bf" +dependencies = [ + "displaydoc", + "yoke", + "zerofrom", +] + +[[package]] +name = "zerovec" +version = "0.11.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "90f911cbc359ab6af17377d242225f4d75119aec87ea711a880987b18cd7b239" +dependencies = [ + "yoke", + "zerofrom", + "zerovec-derive", +] + +[[package]] +name = "zerovec-derive" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "zmij" version = "1.0.21" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 5fb42a6..1888a65 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -2,7 +2,8 @@ name = "rs-chat-api" version = "0.1.0" edition = "2024" -description = "LLM chat application" +description = "RsChat Server" +publish = false [dependencies] anyhow = "1.0.102" @@ -30,7 +31,12 @@ diesel-async = { diesel-jsonb-derive = { path = "./crates/diesel-jsonb-derive" } diesel_migrations = { version = "2.3.2", features = ["postgres"] } dotenvy = "0.15.7" -figment = { version = "0.10.19", features = ["env"] } +figment = { version = "0.10.19", features = ["env", "toml"] } +fred = { + version = "10.1.0", + default-features = false, + features = ["i-keys", "i-streams"] +} hex = "0.4.3" serde = { version = "1.0.228", features = ["derive"] } serde_json = "1.0.150" @@ -59,6 +65,7 @@ tower-sessions = { default-features = false, features = ["axum-core", "memory-store", "private"] } +tower-sessions-redis-store = "0.16.0" tracing = "0.1.44" tracing-appender = "0.2.5" tracing-subscriber = { version = "0.3.23", features = ["env-filter", "json"] } diff --git a/server-new/Dockerfile b/server-new/Dockerfile index 1ebfd5b..676bffc 100644 --- a/server-new/Dockerfile +++ b/server-new/Dockerfile @@ -41,5 +41,6 @@ COPY --from=build --chown=appuser /app/run-server /usr/local/bin/ # Run server WORKDIR /app -ENV RS_CHAT_HOST=0.0.0.0 +COPY --chown=appuser config.toml ./config.toml +ENV RS_CHAT_SERVER__HOST=0.0.0.0 CMD ["run-server"] diff --git a/server-new/config.toml b/server-new/config.toml new file mode 100644 index 0000000..36f99fd --- /dev/null +++ b/server-new/config.toml @@ -0,0 +1,19 @@ +[server] +host = "127.0.0.1" +port = 8080 +log_level = "info" +request_id_header = "x-request-id" + +[database] +url = "postgres://localhost" + +[redis] +url = "redis://localhost:6379" + +[auth] +cookie_name = "auth-rs-chat" +session_length = 604800 + +[security] +body_limit = 2097152 +request_timeout = 120 diff --git a/server-new/src/config.rs b/server-new/src/config.rs index 2214ad7..cd6e85c 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -1,7 +1,8 @@ -use std::net::{IpAddr, Ipv4Addr}; +use std::net::IpAddr; use anyhow::Context; use axum_plugin::AdHocPlugin; +use figment::providers::{Env, Format, Toml}; use serde::Deserialize; use crate::state::AppState; @@ -9,75 +10,60 @@ use crate::state::AppState; /// Parsed app configuration #[derive(Debug, Clone, Deserialize)] pub struct AppConfig { - // Server config - #[serde(default = "default_host")] + pub server: ServerConfig, + pub database: DatabaseConfig, + pub auth: AuthConfig, + pub security: SecurityConfig, + pub redis: RedisConfig, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct ServerConfig { pub host: IpAddr, - #[serde(default = "default_port")] pub port: u16, - #[serde(default = "default_log_level")] pub log_level: String, - #[serde(default = "default_request_id_header")] pub request_id_header: String, pub ip_header: Option, +} - // Database - #[serde(default = "default_database_url")] - pub database_url: String, +#[derive(Debug, Clone, Deserialize)] +pub struct DatabaseConfig { + pub url: String, +} - // Auth +#[derive(Debug, Clone, Deserialize)] +pub struct AuthConfig { pub cookie_key: String, - #[serde(default = "default_cookie_name")] pub cookie_name: String, - #[serde(default = "default_session_length")] pub session_length: i64, +} - // Security - #[serde(default = "default_body_limit")] +#[derive(Debug, Clone, Deserialize)] +pub struct SecurityConfig { pub body_limit: usize, - #[serde(default = "default_req_timeout")] pub request_timeout: u64, } -fn default_host() -> IpAddr { - IpAddr::V4(Ipv4Addr::LOCALHOST) -} -fn default_port() -> u16 { - 8080 -} -fn default_log_level() -> String { - "info".to_string() -} -fn default_request_id_header() -> String { - "x-request-id".to_string() -} -fn default_database_url() -> String { - "postgres://localhost".to_owned() -} -fn default_cookie_name() -> String { - "auth-rs-chat".to_string() -} -fn default_session_length() -> i64 { - 60 * 60 * 24 * 7 // 1 week -} -fn default_body_limit() -> usize { - 2 * 1024 * 1024 // 2 MB -} -fn default_req_timeout() -> u64 { - 120 // 2 minutes + +#[derive(Debug, Clone, Deserialize)] +pub struct RedisConfig { + pub url: String, } /// Plugin that reads and validates configuration, and adds it to server state pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Config").on_init(async |mut state| { let config = extract_config()?; + tracing::info!(log_level = config.server.log_level, "Config loaded!"); state.insert(config); Ok(state) }) } -/// Extract the configuration from env variables prefixed with `RS_CHAT_`. +/// Extract configuration from config.toml, then environment overrides. fn extract_config() -> anyhow::Result { let config = figment::Figment::new() - .merge(figment::providers::Env::prefixed("RS_CHAT_")) + .merge(Toml::file("config.toml")) + .merge(Env::prefixed("RS_CHAT_").split("__")) .extract::() .context("Failed to extract valid configuration")?; diff --git a/server-new/src/extractors/session.rs b/server-new/src/extractors/session.rs index d1157d5..9888b06 100644 --- a/server-new/src/extractors/session.rs +++ b/server-new/src/extractors/session.rs @@ -79,6 +79,7 @@ impl FromRequestParts for SessionMeta { ) -> Result { let ip_header = state .config + .server .ip_header .as_ref() .and_then(|h| parts.headers.get(h).and_then(|h| h.to_str().ok())) diff --git a/server-new/src/lib.rs b/server-new/src/lib.rs index 38ae43f..17895f9 100644 --- a/server-new/src/lib.rs +++ b/server-new/src/lib.rs @@ -15,6 +15,7 @@ pub async fn create_app() -> anyhow::Result> { let app = App::new() .register(config::plugin()) // Extract configuration and add to state .register(plugins::database::plugin()) // Initialize database + .register(plugins::redis::plugin()) // Initialize Redis .register(routes::plugin()) // Add API routes .register(plugins::session::plugin()) // Setup sessions .register(plugins::logging::plugin()) // Request logging diff --git a/server-new/src/main.rs b/server-new/src/main.rs index 655290e..fb2a83a 100644 --- a/server-new/src/main.rs +++ b/server-new/src/main.rs @@ -20,11 +20,11 @@ async fn main() -> anyhow::Result<()> { // Set log level from config let env_filter = tracing_subscriber::EnvFilter::builder() .with_default_directive(LevelFilter::INFO.into()) - .parse(&config.log_level)?; + .parse(&config.server.log_level)?; log_filter_handle.reload(env_filter)?; // Start listening for requests - let addr = SocketAddr::new(config.host, config.port); + let addr = SocketAddr::new(config.server.host, config.server.port); let listener = tokio::net::TcpListener::bind(addr).await?; tracing::info!("Server listening on http://{}...", listener.local_addr()?); axum::serve( @@ -42,7 +42,9 @@ fn init_logging() -> ( tracing_subscriber::reload::Handle, tracing_appender::non_blocking::WorkerGuard, ) { - let init_log_level = std::env::var("RS_CHAT_LOG_LEVEL").unwrap_or("info".into()); + let init_log_level = std::env::var("RS_CHAT_SERVER__LOG_LEVEL") + .or_else(|_| std::env::var("RS_CHAT_LOG_LEVEL")) + .unwrap_or("info".into()); let (writer, guard) = tracing_appender::non_blocking(std::io::stdout()); let (filter_layer, filter_handle) = tracing_subscriber::reload::Layer::new(EnvFilter::new(init_log_level)); diff --git a/server-new/src/plugins/database.rs b/server-new/src/plugins/database.rs index 51e992d..d1fe098 100644 --- a/server-new/src/plugins/database.rs +++ b/server-new/src/plugins/database.rs @@ -16,10 +16,11 @@ pub fn plugin() -> AdHocPlugin { let app_config = state.get::().context("missing config")?; let manager = - AsyncDieselConnectionManager::::new(&app_config.database_url); + AsyncDieselConnectionManager::::new(&app_config.database.url); let pool: DbPool = Pool::builder(manager).build()?; let cxn = pool.get().await.context("failed to connect to database")?; + tracing::info!("Connected to database"); match AsyncMigrationHarness::new(cxn).run_pending_migrations(MIGRATIONS) { Ok(run_migrations) if run_migrations.is_empty() => { tracing::info!("No migrations to run"); diff --git a/server-new/src/plugins/logging.rs b/server-new/src/plugins/logging.rs index f9fc373..a821105 100644 --- a/server-new/src/plugins/logging.rs +++ b/server-new/src/plugins/logging.rs @@ -15,7 +15,7 @@ use crate::state::AppState; pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Request logs").on_setup(|router, state: &AppState| { const LOG_LEVEL: Level = Level::INFO; - let request_id_header = HeaderName::from_str(&state.config.request_id_header) + let request_id_header = HeaderName::from_str(&state.config.server.request_id_header) .context("invalid request ID header")?; let trace_layer = TraceLayer::new_for_http() diff --git a/server-new/src/plugins/mod.rs b/server-new/src/plugins/mod.rs index e439465..dfe31c0 100644 --- a/server-new/src/plugins/mod.rs +++ b/server-new/src/plugins/mod.rs @@ -1,4 +1,5 @@ pub mod database; pub mod logging; +pub mod redis; pub mod security; pub mod session; diff --git a/server-new/src/plugins/redis.rs b/server-new/src/plugins/redis.rs new file mode 100644 index 0000000..7be9cca --- /dev/null +++ b/server-new/src/plugins/redis.rs @@ -0,0 +1,33 @@ +use std::time::Duration; + +use anyhow::Context; +use axum_plugin::AdHocPlugin; +use fred::prelude::*; + +use crate::{config::AppConfig, db::DbPool, state::AppState}; + +const DEFAULT_TIMEOUT: Duration = Duration::from_secs(8); + +pub fn plugin() -> AdHocPlugin { + AdHocPlugin::named("Redis").on_init(async |mut state| { + let app_config = state.get::().context("no config")?; + let db_pool = state.get::().context("no database pool")?; + let config = Config::from_url(&app_config.redis.url).context("invalid Redis URL")?; + let pool = Builder::from_config(config) + .with_connection_config(|c| { + c.connection_timeout = DEFAULT_TIMEOUT; + c.internal_command_timeout = DEFAULT_TIMEOUT; + c.tcp.nodelay = Some(true); + }) + .with_performance_config(|c| { + c.default_command_timeout = DEFAULT_TIMEOUT; + }) + .build_pool(db_pool.status().max_size)?; // same size as database pool + + pool.init().await.context("failed to connect to Redis")?; + tracing::info!("Connected to Redis"); + + state.insert(pool); + Ok(state) + }) +} diff --git a/server-new/src/plugins/security.rs b/server-new/src/plugins/security.rs index 9b7f0d9..d48917b 100644 --- a/server-new/src/plugins/security.rs +++ b/server-new/src/plugins/security.rs @@ -20,10 +20,10 @@ pub fn plugin() -> AdHocPlugin { .into_layer()?; let service = ServiceBuilder::new() - .layer(RequestBodyLimitLayer::new(state.config.body_limit)) + .layer(RequestBodyLimitLayer::new(state.config.security.body_limit)) .layer(TimeoutLayer::with_status_code( StatusCode::REQUEST_TIMEOUT, - Duration::from_secs(state.config.request_timeout), + Duration::from_secs(state.config.security.request_timeout), )) .layer(security_headers); diff --git a/server-new/src/plugins/session.rs b/server-new/src/plugins/session.rs index f86e32c..4d5ec80 100644 --- a/server-new/src/plugins/session.rs +++ b/server-new/src/plugins/session.rs @@ -10,16 +10,16 @@ use crate::{services::SessionDbStore, state::AppState}; pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Session").on_setup(|router, state: &AppState| { let cookie_key = - hex::decode(&state.config.cookie_key).context("cookie_key must be hex value")?; + hex::decode(&state.config.auth.cookie_key).context("cookie_key must be hex value")?; if cookie_key.len() < 32 { bail!("cookie_key must be at least 32 bytes"); } let session_store = SessionDbStore::new(state.db_pool.clone()); let session_layer = SessionManagerLayer::new(session_store) - .with_name(state.config.cookie_name.clone()) + .with_name(state.config.auth.cookie_name.clone()) .with_expiry(Expiry::OnInactivity(Duration::seconds( - state.config.session_length, + state.config.auth.session_length, ))) .with_private(Key::derive_from(&cookie_key)) .with_path("/") diff --git a/server-new/src/state.rs b/server-new/src/state.rs index f62f7cd..fd7b37d 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -14,6 +14,7 @@ pub struct AppState(Arc); pub struct AppStateInner { pub config: AppConfig, pub db_pool: DbPool, + pub redis: fred::prelude::Pool, } impl AppState { From 7a13e02089169920640fcd69b399f7d68a2243f2 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 01:43:18 -0400 Subject: [PATCH 028/111] cache sessions in redis --- server-new/Cargo.lock | 30 ++++-------------------------- server-new/Cargo.toml | 12 +++++------- server-new/src/plugins/session.rs | 10 ++++++++-- 3 files changed, 17 insertions(+), 35 deletions(-) diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index edca47d..f7a2bce 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -1167,7 +1167,6 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "224399e74b87b5f3557511d98dff8b14089b3dadafcab6bb93eab67d3aace965" dependencies = [ "scopeguard", - "serde", ] [[package]] @@ -2243,31 +2242,11 @@ dependencies = [ "tower-cookies", "tower-layer", "tower-service", - "tower-sessions-core 0.15.0", + "tower-sessions-core", "tower-sessions-memory-store", "tracing", ] -[[package]] -name = "tower-sessions-core" -version = "0.14.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ce8cce604865576b7751b7a6bc3058f754569a60d689328bb74c52b1d87e355b" -dependencies = [ - "async-trait", - "base64", - "futures", - "http", - "parking_lot", - "rand 0.8.6", - "serde", - "serde_json", - "thiserror", - "time", - "tokio", - "tracing", -] - [[package]] name = "tower-sessions-core" version = "0.15.0" @@ -2298,21 +2277,20 @@ dependencies = [ "async-trait", "time", "tokio", - "tower-sessions-core 0.15.0", + "tower-sessions-core", ] [[package]] name = "tower-sessions-redis-store" version = "0.16.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8e15b774f3d46625a27a8ac1238ecd73c8bd50013244e2de004026e161aad728" +source = "git+https://github.com/maxcountryman/tower-sessions-stores?rev=69e025f#69e025f97b8b6ca54618e000375cb1aaa852209a" dependencies = [ "async-trait", "fred", "rmp-serde", "thiserror", "time", - "tower-sessions-core 0.14.0", + "tower-sessions-core", ] [[package]] diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 1888a65..1fa40dd 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -53,19 +53,17 @@ tokio = { tower = { version = "0.5", default-features = false } tower-http = { version = "0.7.0", - features = [ - "limit", - "request-id", - "timeout", - "trace", - ] + features = ["limit", "request-id", "timeout", "trace"] } tower-sessions = { version = "0.15.0", default-features = false, features = ["axum-core", "memory-store", "private"] } -tower-sessions-redis-store = "0.16.0" +tower-sessions-redis-store = { + git = "https://github.com/maxcountryman/tower-sessions-stores", + rev = "69e025f" +} tracing = "0.1.44" tracing-appender = "0.2.5" tracing-subscriber = { version = "0.3.23", features = ["env-filter", "json"] } diff --git a/server-new/src/plugins/session.rs b/server-new/src/plugins/session.rs index 4d5ec80..1a8654d 100644 --- a/server-new/src/plugins/session.rs +++ b/server-new/src/plugins/session.rs @@ -1,12 +1,16 @@ use anyhow::{Context, bail}; use axum_plugin::AdHocPlugin; use tower_sessions::{ - Expiry, SessionManagerLayer, + CachingSessionStore, Expiry, SessionManagerLayer, cookie::{Key, SameSite, time::Duration}, }; +use tower_sessions_redis_store::RedisStore; use crate::{services::SessionDbStore, state::AppState}; +const REDIS_PREFIX: &str = "rs-chat:sess:"; + +/// Add session handling to the server. Sessions are stored in Postgres and cached in Redis. pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Session").on_setup(|router, state: &AppState| { let cookie_key = @@ -15,7 +19,9 @@ pub fn plugin() -> AdHocPlugin { bail!("cookie_key must be at least 32 bytes"); } - let session_store = SessionDbStore::new(state.db_pool.clone()); + let redis_store = RedisStore::with_prefix(state.redis.clone(), REDIS_PREFIX.to_owned()); + let db_store = SessionDbStore::new(state.db_pool.clone()); + let session_store = CachingSessionStore::new(redis_store, db_store); let session_layer = SessionManagerLayer::new(session_store) .with_name(state.config.auth.cookie_name.clone()) .with_expiry(Expiry::OnInactivity(Duration::seconds( From b651198785acb21dcef70b4d6ffaf4aada08e3fd Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 01:56:27 -0400 Subject: [PATCH 029/111] redis shutdown --- server-new/bacon.toml | 6 +++ .../{repositories/mod.rs => repositories.rs} | 0 server-new/src/db/repositories/user.rs | 18 +++---- server-new/src/plugins/redis.rs | 50 ++++++++++++------- 4 files changed, 46 insertions(+), 28 deletions(-) create mode 100644 server-new/bacon.toml rename server-new/src/db/{repositories/mod.rs => repositories.rs} (100%) diff --git a/server-new/bacon.toml b/server-new/bacon.toml new file mode 100644 index 0000000..8d5e596 --- /dev/null +++ b/server-new/bacon.toml @@ -0,0 +1,6 @@ +[jobs.dev] +command = ["cargo", "run"] +need_stdout = true +background = false +on_change_strategy = "kill_then_restart" +kill = ["kill", "-s", "INT"] diff --git a/server-new/src/db/repositories/mod.rs b/server-new/src/db/repositories.rs similarity index 100% rename from server-new/src/db/repositories/mod.rs rename to server-new/src/db/repositories.rs diff --git a/server-new/src/db/repositories/user.rs b/server-new/src/db/repositories/user.rs index 8968284..444ab1a 100644 --- a/server-new/src/db/repositories/user.rs +++ b/server-new/src/db/repositories/user.rs @@ -25,16 +25,16 @@ impl<'a> UserRepository<'a> { Ok(user) } - // pub async fn find_by_github_id(&mut self, id: &str) -> Result, Error> { - // let user = users::table - // .filter(users::github_id.eq(id)) - // .select(ChatRsUser::as_select()) - // .first(self.db) - // .await - // .optional()?; + pub async fn find_by_github_id(&mut self, id: &str) -> Result, Error> { + let user = users::table + .filter(users::github_id.eq(id)) + .select(ChatRsUser::as_select()) + .first(self.db) + .await + .optional()?; - // Ok(user) - // } + Ok(user) + } // pub async fn find_by_google_id(&mut self, id: &str) -> Result, Error> { // let user = users::table diff --git a/server-new/src/plugins/redis.rs b/server-new/src/plugins/redis.rs index 7be9cca..ce4d1c1 100644 --- a/server-new/src/plugins/redis.rs +++ b/server-new/src/plugins/redis.rs @@ -9,25 +9,37 @@ use crate::{config::AppConfig, db::DbPool, state::AppState}; const DEFAULT_TIMEOUT: Duration = Duration::from_secs(8); pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("Redis").on_init(async |mut state| { - let app_config = state.get::().context("no config")?; - let db_pool = state.get::().context("no database pool")?; - let config = Config::from_url(&app_config.redis.url).context("invalid Redis URL")?; - let pool = Builder::from_config(config) - .with_connection_config(|c| { - c.connection_timeout = DEFAULT_TIMEOUT; - c.internal_command_timeout = DEFAULT_TIMEOUT; - c.tcp.nodelay = Some(true); - }) - .with_performance_config(|c| { - c.default_command_timeout = DEFAULT_TIMEOUT; - }) - .build_pool(db_pool.status().max_size)?; // same size as database pool + AdHocPlugin::named("Redis") + .on_init(async |mut state| { + let app_config = state.get::().context("no config")?; + let db_pool = state.get::().context("no database pool")?; + let config = Config::from_url(&app_config.redis.url).context("invalid Redis URL")?; + let pool = Builder::from_config(config) + .with_connection_config(|c| { + c.connection_timeout = DEFAULT_TIMEOUT; + c.internal_command_timeout = DEFAULT_TIMEOUT; + c.tcp.nodelay = Some(true); + }) + .with_performance_config(|c| { + c.default_command_timeout = DEFAULT_TIMEOUT; + }) + .build_pool(db_pool.status().max_size)?; // same size as database pool - pool.init().await.context("failed to connect to Redis")?; - tracing::info!("Connected to Redis"); + pool.init().await.context("failed to connect to Redis")?; + tracing::info!("Connected to Redis"); - state.insert(pool); - Ok(state) - }) + state.insert(pool); + Ok(state) + }) + .on_shutdown(|state: &AppState| { + let pool = state.redis.clone(); + async move { + if let Err(e) = pool.quit().await { + tracing::warn!("Error shutting down Redis pool: {e}"); + } else { + tracing::info!("Shut down Redis pool") + } + Ok(()) + } + }) } From 19caaf4541b1dc479ecd1dfc3cd44be96fb9fb30 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 01:58:48 -0400 Subject: [PATCH 030/111] rename /routes to /api --- server-new/src/{routes => api}/auth.rs | 0 server-new/src/{routes => api}/health.rs | 0 server-new/src/{routes => api}/hello.rs | 0 server-new/src/{routes => api}/mod.rs | 0 server-new/src/lib.rs | 4 ++-- 5 files changed, 2 insertions(+), 2 deletions(-) rename server-new/src/{routes => api}/auth.rs (100%) rename server-new/src/{routes => api}/health.rs (100%) rename server-new/src/{routes => api}/hello.rs (100%) rename server-new/src/{routes => api}/mod.rs (100%) diff --git a/server-new/src/routes/auth.rs b/server-new/src/api/auth.rs similarity index 100% rename from server-new/src/routes/auth.rs rename to server-new/src/api/auth.rs diff --git a/server-new/src/routes/health.rs b/server-new/src/api/health.rs similarity index 100% rename from server-new/src/routes/health.rs rename to server-new/src/api/health.rs diff --git a/server-new/src/routes/hello.rs b/server-new/src/api/hello.rs similarity index 100% rename from server-new/src/routes/hello.rs rename to server-new/src/api/hello.rs diff --git a/server-new/src/routes/mod.rs b/server-new/src/api/mod.rs similarity index 100% rename from server-new/src/routes/mod.rs rename to server-new/src/api/mod.rs diff --git a/server-new/src/lib.rs b/server-new/src/lib.rs index 17895f9..f89cde8 100644 --- a/server-new/src/lib.rs +++ b/server-new/src/lib.rs @@ -2,12 +2,12 @@ use axum_plugin::{App, InitializedApp}; use crate::state::AppState; +mod api; mod config; mod db; mod error; mod extractors; mod plugins; -mod routes; mod services; mod state; @@ -16,7 +16,7 @@ pub async fn create_app() -> anyhow::Result> { .register(config::plugin()) // Extract configuration and add to state .register(plugins::database::plugin()) // Initialize database .register(plugins::redis::plugin()) // Initialize Redis - .register(routes::plugin()) // Add API routes + .register(api::plugin()) // Add API routes .register(plugins::session::plugin()) // Setup sessions .register(plugins::logging::plugin()) // Request logging .register(plugins::security::plugin()) // Body limit, security headers, etc. From da61e5e7f56e2c4f15143ef9ba7e8bf92a3dbdf6 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 17:36:46 -0400 Subject: [PATCH 031/111] oauth --- server-new/.env.example | 5 +- server-new/Cargo.lock | 780 +++++++++++++++++- server-new/Cargo.toml | 4 + server-new/src/api/auth.rs | 72 +- server-new/src/api/mod.rs | 11 +- server-new/src/config.rs | 8 +- server-new/src/db/repositories/user.rs | 66 +- server-new/src/error.rs | 8 +- server-new/src/extractors/session.rs | 2 +- server-new/src/lib.rs | 4 + server-new/src/main.rs | 4 +- server-new/src/plugins/session.rs | 9 +- server-new/src/services/auth/mod.rs | 35 +- server-new/src/services/auth/oauth/discord.rs | 108 +++ server-new/src/services/auth/oauth/github.rs | 109 +++ server-new/src/services/auth/oauth/mod.rs | 228 +++++ server-new/src/services/mod.rs | 4 +- server-new/src/state.rs | 5 +- 18 files changed, 1369 insertions(+), 93 deletions(-) create mode 100644 server-new/src/services/auth/oauth/discord.rs create mode 100644 server-new/src/services/auth/oauth/github.rs create mode 100644 server-new/src/services/auth/oauth/mod.rs diff --git a/server-new/.env.example b/server-new/.env.example index a863a03..9660a89 100644 --- a/server-new/.env.example +++ b/server-new/.env.example @@ -1,12 +1,11 @@ -# Server Configuration -RS_CHAT_SERVER__PORT=8080 - # Database (2nd one is needed for Diesel CLI) RS_CHAT_DATABASE__URL=postgres://postgres:postgres@localhost/postgres DATABASE_URL=postgres://postgres:postgres@localhost/postgres # Auth RS_CHAT_AUTH__COOKIE_KEY= # hex secret >=32 bytes, e.g. `openssl rand --hex 32` +RS_CHAT_AUTH__GITHUB__CLIENT_ID= +RS_CHAT_AUTH__GITHUB__CLIENT_SECRET= # Logging RS_CHAT_SERVER__LOG_LEVEL=info diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index f7a2bce..42fd158 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -8,7 +8,7 @@ version = "0.5.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0" dependencies = [ - "crypto-common", + "crypto-common 0.1.7", "generic-array", ] @@ -20,7 +20,7 @@ checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0" dependencies = [ "cfg-if", "cipher", - "cpufeatures", + "cpufeatures 0.2.17", ] [[package]] @@ -70,6 +70,24 @@ dependencies = [ "rustversion", ] +[[package]] +name = "async-oauth2" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b82d800aa8c98755d56f8492f8861f6f958e750e1ce940f55846e94d4defef28" +dependencies = [ + "base64", + "bytes", + "http", + "rand 0.10.1", + "reqwest", + "serde", + "serde-aux", + "serde_json", + "sha2 0.11.0", + "url", +] + [[package]] name = "async-trait" version = "0.1.89" @@ -102,6 +120,28 @@ version = "1.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" +[[package]] +name = "aws-lc-rs" +version = "1.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5ec2f1fc3ec205783a5da9a7e6c1509cc69dedf09a1949e412c1e18469326d00" +dependencies = [ + "aws-lc-sys", + "zeroize", +] + +[[package]] +name = "aws-lc-sys" +version = "0.41.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1a2f9779ce85b93ab6170dd940ad0169b5766ff848247aff13bb788b832fe3f4" +dependencies = [ + "cc", + "cmake", + "dunce", + "fs_extra", +] + [[package]] name = "axum" version = "0.8.9" @@ -210,6 +250,15 @@ dependencies = [ "generic-array", ] +[[package]] +name = "block-buffer" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2f6c7dbe95a6ed67ad9f18e57daf93a2f034c524b99fd2b76d18fdfeb6660aa" +dependencies = [ + "hybrid-array", +] + [[package]] name = "bumpalo" version = "3.20.3" @@ -251,6 +300,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e228eec9be7c17ccb640b59b36a5cd805ea2a564a4c5e162c2f659fea30d3b96" dependencies = [ "find-msvc-tools", + "jobserver", + "libc", "shlex", ] @@ -260,6 +311,23 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" +[[package]] +name = "cfg_aliases" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" + +[[package]] +name = "chacha20" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6f8d983286843e49675a4b7a2d174efe136dc93a18d69130dd18198a6c167601" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "rand_core 0.10.1", +] + [[package]] name = "chrono" version = "0.4.45" @@ -278,10 +346,35 @@ version = "0.4.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" dependencies = [ - "crypto-common", + "crypto-common 0.1.7", "inout", ] +[[package]] +name = "cmake" +version = "0.1.58" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c0f78a02292a74a88ac736019ab962ece0bc380e3f977bf72e376c5d78ff0678" +dependencies = [ + "cc", +] + +[[package]] +name = "combine" +version = "4.6.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba5a308b75df32fe02788e748662718f03fde005016435c444eea572398219fd" +dependencies = [ + "bytes", + "memchr", +] + +[[package]] +name = "const-oid" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6ef517f0926dd24a1582492c791b6a4818a4d94e789a334894aa15b0d12f55c" + [[package]] name = "cookie" version = "0.18.1" @@ -294,7 +387,7 @@ dependencies = [ "hmac", "percent-encoding", "rand 0.8.6", - "sha2", + "sha2 0.10.9", "subtle", "time", "version_check", @@ -306,6 +399,26 @@ version = "0.3.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "396de984970346b0d9e93d1415082923c679e5ae5c3ee3dcbd104f5610af126b" +[[package]] +name = "core-foundation" +version = "0.9.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91e195e091a93c46f7102ec7818a2aa394e1e1771c3ab4825963fa03e45afb8f" +dependencies = [ + "core-foundation-sys", + "libc", +] + +[[package]] +name = "core-foundation" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b2a6cd9ae233e7f62ba4e9353e81a88df7fc8a5987b8d445b4d90c879bd156f6" +dependencies = [ + "core-foundation-sys", + "libc", +] + [[package]] name = "core-foundation-sys" version = "0.8.7" @@ -321,6 +434,15 @@ dependencies = [ "libc", ] +[[package]] +name = "cpufeatures" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" +dependencies = [ + "libc", +] + [[package]] name = "crc16" version = "0.4.0" @@ -353,6 +475,15 @@ dependencies = [ "typenum", ] +[[package]] +name = "crypto-common" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ce6e4c961d6cd6c9a86db418387425e8bdeaf05b3c8bc1411e6dca4c252f1453" +dependencies = [ + "hybrid-array", +] + [[package]] name = "ctr" version = "0.9.2" @@ -537,11 +668,22 @@ version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ - "block-buffer", - "crypto-common", + "block-buffer 0.10.4", + "crypto-common 0.1.7", "subtle", ] +[[package]] +name = "digest" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1dd6dbb5841937940781866fa1281a1ff7bd3bf827091440879f9994983d5c2" +dependencies = [ + "block-buffer 0.12.1", + "const-oid", + "crypto-common 0.2.2", +] + [[package]] name = "displaydoc" version = "0.2.6" @@ -579,12 +721,27 @@ dependencies = [ "syn", ] +[[package]] +name = "dunce" +version = "1.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92773504d58c093f6de2459af4af33faa518c13451eb8f2b5698ed3d36e7c813" + [[package]] name = "either" version = "1.16.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" +[[package]] +name = "encoding_rs" +version = "0.8.35" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75030f3c4f45dafd7586dd6780965a8c7e8e285a5ecb86713e63a79c5b2766f3" +dependencies = [ + "cfg-if", +] + [[package]] name = "equivalent" version = "1.0.2" @@ -688,6 +845,12 @@ dependencies = [ "syn", ] +[[package]] +name = "fs_extra" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" + [[package]] name = "futures" version = "0.3.32" @@ -793,8 +956,10 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0" dependencies = [ "cfg-if", + "js-sys", "libc", "wasi 0.11.1+wasi-snapshot-preview1", + "wasm-bindgen", ] [[package]] @@ -804,9 +969,11 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" dependencies = [ "cfg-if", + "js-sys", "libc", "r-efi 5.3.0", "wasip2", + "wasm-bindgen", ] [[package]] @@ -818,6 +985,7 @@ dependencies = [ "cfg-if", "libc", "r-efi 6.0.0", + "rand_core 0.10.1", ] [[package]] @@ -830,6 +998,25 @@ dependencies = [ "polyval", ] +[[package]] +name = "h2" +version = "0.4.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6cb093c84e8bd9b188d4c4a8cb6579fc016968d14c99882163cd3ff402a4f155" +dependencies = [ + "atomic-waker", + "bytes", + "fnv", + "futures-core", + "futures-sink", + "http", + "indexmap", + "slab", + "tokio", + "tokio-util", + "tracing", +] + [[package]] name = "hashbrown" version = "0.17.1" @@ -875,7 +1062,7 @@ version = "0.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" dependencies = [ - "digest", + "digest 0.10.7", ] [[package]] @@ -923,6 +1110,15 @@ version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" +[[package]] +name = "hybrid-array" +version = "0.4.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9155a582abd142abc056962c29e3ce5ff2ad5469f4246b537ed42c5deba857da" +dependencies = [ + "typenum", +] + [[package]] name = "hyper" version = "1.10.1" @@ -933,6 +1129,7 @@ dependencies = [ "bytes", "futures-channel", "futures-core", + "h2", "http", "http-body", "httparse", @@ -941,6 +1138,22 @@ dependencies = [ "pin-project-lite", "smallvec", "tokio", + "want", +] + +[[package]] +name = "hyper-rustls" +version = "0.27.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "33ca68d021ef39cf6463ab54c1d0f5daf03377b70561305bb89a8f83aab66e0f" +dependencies = [ + "http", + "hyper", + "hyper-util", + "rustls", + "tokio", + "tokio-rustls", + "tower-service", ] [[package]] @@ -949,13 +1162,23 @@ version = "0.1.20" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0" dependencies = [ + "base64", "bytes", + "futures-channel", + "futures-util", "http", "http-body", "hyper", + "ipnet", + "libc", + "percent-encoding", "pin-project-lite", + "socket2 0.6.4", + "system-configuration", "tokio", "tower-service", + "tracing", + "windows-registry", ] [[package]] @@ -1116,12 +1339,77 @@ dependencies = [ "generic-array", ] +[[package]] +name = "ipnet" +version = "2.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d98f6fed1fde3f8c21bc40a1abb88dd75e67924f9cffc3ef95607bad8017f8e2" + [[package]] name = "itoa" version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" +[[package]] +name = "jni" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5efd9a482cf3a427f00d6b35f14332adc7902ce91efb778580e180ff90fa3498" +dependencies = [ + "cfg-if", + "combine", + "jni-macros", + "jni-sys", + "log", + "simd_cesu8", + "thiserror", + "walkdir", + "windows-link", +] + +[[package]] +name = "jni-macros" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a00109accc170f0bdb141fed3e393c565b6f5e072365c3bd58f5b062591560a3" +dependencies = [ + "proc-macro2", + "quote", + "rustc_version", + "simd_cesu8", + "syn", +] + +[[package]] +name = "jni-sys" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6377a88cb3910bee9b0fa88d4f42e1d2da8e79915598f65fb0c7ee14c878af2" +dependencies = [ + "jni-sys-macros", +] + +[[package]] +name = "jni-sys-macros" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264" +dependencies = [ + "quote", + "syn", +] + +[[package]] +name = "jobserver" +version = "0.1.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9afb3de4395d6b3e67a780b6de64b51c978ecf11cb9a462c66be7d4ca9039d33" +dependencies = [ + "getrandom 0.3.4", + "libc", +] + [[package]] name = "js-sys" version = "0.3.102" @@ -1175,6 +1463,12 @@ version = "0.4.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0ceec5bc11778974d1bcb055b18002eba7f4b3518b6a0081b3af5f21666da9ad" +[[package]] +name = "lru-slab" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "112b39cec0b298b6c1999fee3e31427f74f676e4cb9879ed1a121b43661a4154" + [[package]] name = "matchers" version = "0.2.0" @@ -1197,7 +1491,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d89e7ee0cfbedfc4da3340218492196241d89eefb6dab27de5df917a6d2e78cf" dependencies = [ "cfg-if", - "digest", + "digest 0.10.7", ] [[package]] @@ -1324,6 +1618,21 @@ version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" +[[package]] +name = "openssl-probe" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe" + +[[package]] +name = "ordered-float" +version = "2.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68f19d67e5a2795c94e73e0bb1cc1a7edeb2e28efd39e2e1c9b7a40c1108b11c" +dependencies = [ + "num-traits", +] + [[package]] name = "parking_lot" version = "0.12.5" @@ -1408,7 +1717,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.2.17", "opaque-debug", "universal-hash", ] @@ -1427,7 +1736,7 @@ dependencies = [ "md-5", "memchr", "rand 0.9.4", - "sha2", + "sha2 0.10.9", "stringprep", ] @@ -1488,6 +1797,62 @@ dependencies = [ "yansi", ] +[[package]] +name = "quinn" +version = "0.11.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c1a41e437b6bbd489372cd4971de128e85c855f56c57f283d20ff016cf7c0a8" +dependencies = [ + "bytes", + "cfg_aliases", + "pin-project-lite", + "quinn-proto", + "quinn-udp", + "rustc-hash", + "rustls", + "socket2 0.6.4", + "thiserror", + "tokio", + "tracing", + "web-time", +] + +[[package]] +name = "quinn-proto" +version = "0.11.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4fcb935c5bec503c2f0e306bdd3e58bb9029dcb14fa8d9ac76e3a5256ac0763e" +dependencies = [ + "aws-lc-rs", + "bytes", + "getrandom 0.3.4", + "lru-slab", + "rand 0.9.4", + "ring", + "rustc-hash", + "rustls", + "rustls-pki-types", + "slab", + "thiserror", + "tinyvec", + "tracing", + "web-time", +] + +[[package]] +name = "quinn-udp" +version = "0.5.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "addec6a0dcad8a8d96a771f815f0eaf55f9d1805756410b39f5fa81332574cbd" +dependencies = [ + "cfg_aliases", + "libc", + "once_cell", + "socket2 0.6.4", + "tracing", + "windows-sys 0.52.0", +] + [[package]] name = "quote" version = "1.0.46" @@ -1530,6 +1895,17 @@ dependencies = [ "rand_core 0.9.5", ] +[[package]] +name = "rand" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2e8e8bcc7961af1fdac401278c6a831614941f6164ee3bf4ce61b7edb162207" +dependencies = [ + "chacha20", + "getrandom 0.4.3", + "rand_core 0.10.1", +] + [[package]] name = "rand_chacha" version = "0.3.1" @@ -1568,6 +1944,12 @@ dependencies = [ "getrandom 0.3.4", ] +[[package]] +name = "rand_core" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" + [[package]] name = "redis-protocol" version = "6.0.0" @@ -1608,6 +1990,60 @@ version = "0.8.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" +[[package]] +name = "reqwest" +version = "0.13.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "219c5811de6525e5416c7d5d53bb656d3afdbc6c5af816e0802bcfa42dbdc1c3" +dependencies = [ + "base64", + "bytes", + "encoding_rs", + "futures-core", + "h2", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-rustls", + "hyper-util", + "js-sys", + "log", + "mime", + "percent-encoding", + "pin-project-lite", + "quinn", + "rustls", + "rustls-pki-types", + "rustls-platform-verifier", + "serde", + "serde_json", + "sync_wrapper", + "tokio", + "tokio-rustls", + "tower", + "tower-http 0.6.11", + "tower-service", + "url", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + +[[package]] +name = "ring" +version = "0.17.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4689e6c2294d81e88dc6261c768b63bc4fcdb852be6d1352498b114f61383b7" +dependencies = [ + "cc", + "cfg-if", + "getrandom 0.2.17", + "libc", + "untrusted", + "windows-sys 0.52.0", +] + [[package]] name = "rmp" version = "0.8.15" @@ -1632,6 +2068,7 @@ name = "rs-chat-api" version = "0.1.0" dependencies = [ "anyhow", + "async-oauth2", "async-trait", "axum", "axum-helmet", @@ -1644,13 +2081,16 @@ dependencies = [ "dotenvy", "figment", "fred", + "futures", "hex", + "reqwest", "serde", "serde_json", "serde_with", + "subtle", "tokio", "tower", - "tower-http", + "tower-http 0.7.0", "tower-sessions", "tower-sessions-redis-store", "tracing", @@ -1665,6 +2105,90 @@ version = "2.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "94300abf3f1ae2e2b8ffb7b58043de3d399c73fa6f4b73826402a5c457614dbe" +[[package]] +name = "rustc_version" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfcb3a22ef46e85b45de6ee7e79d063319ebb6594faafcf1c225ea92ab6e9b92" +dependencies = [ + "semver", +] + +[[package]] +name = "rustls" +version = "0.23.41" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6b92b125634d9b795e7beca796cc790df15a7fb38323bf3196fda83292d06b1f" +dependencies = [ + "aws-lc-rs", + "once_cell", + "rustls-pki-types", + "rustls-webpki", + "subtle", + "zeroize", +] + +[[package]] +name = "rustls-native-certs" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dab5152771c58876a2146916e53e35057e1a4dfa2b9df0f0305b07f611fdea4d" +dependencies = [ + "openssl-probe", + "rustls-pki-types", + "schannel", + "security-framework", +] + +[[package]] +name = "rustls-pki-types" +version = "1.14.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "30a7197ae7eb376e574fe940d068c30fe0462554a3ddbe4eca7838e049c937a9" +dependencies = [ + "web-time", + "zeroize", +] + +[[package]] +name = "rustls-platform-verifier" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "26d1e2536ce4f35f4846aa13bff16bd0ff40157cdb14cc056c7b14ba41233ba0" +dependencies = [ + "core-foundation 0.10.1", + "core-foundation-sys", + "jni", + "log", + "once_cell", + "rustls", + "rustls-native-certs", + "rustls-platform-verifier-android", + "rustls-webpki", + "security-framework", + "security-framework-sys", + "webpki-root-certs", + "windows-sys 0.61.2", +] + +[[package]] +name = "rustls-platform-verifier-android" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f87165f0995f63a9fbeea62b64d10b4d9d8e78ec6d7d51fb2125fda7bb36788f" + +[[package]] +name = "rustls-webpki" +version = "0.103.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" +dependencies = [ + "aws-lc-rs", + "ring", + "rustls-pki-types", + "untrusted", +] + [[package]] name = "rustversion" version = "1.0.22" @@ -1677,12 +2201,53 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" +[[package]] +name = "same-file" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93fc1dc3aaa9bfed95e02e6eadabb4baf7e3078b0bd1b4d7b6b0b68378900502" +dependencies = [ + "winapi-util", +] + +[[package]] +name = "schannel" +version = "0.1.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91c1b7e4904c873ef0710c1f407dde2e6287de2bebc1bbbf7d430bb7cbffd939" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "scopeguard" version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" +[[package]] +name = "security-framework" +version = "3.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d" +dependencies = [ + "bitflags", + "core-foundation 0.10.1", + "core-foundation-sys", + "libc", + "security-framework-sys", +] + +[[package]] +name = "security-framework-sys" +version = "2.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce2691df843ecc5d231c0b14ece2acc3efb62c0a398c7e1d875f3983ce020e3" +dependencies = [ + "core-foundation-sys", + "libc", +] + [[package]] name = "semver" version = "1.0.28" @@ -1699,6 +2264,28 @@ dependencies = [ "serde_derive", ] +[[package]] +name = "serde-aux" +version = "4.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "207f67b28fe90fb596503a9bf0bf1ea5e831e21307658e177c5dfcdfc3ab8a0a" +dependencies = [ + "chrono", + "serde", + "serde-value", + "serde_json", +] + +[[package]] +name = "serde-value" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f3a1a3341211875ef120e117ea7fd5228530ae7e7036a779fdc9117be6b3282c" +dependencies = [ + "ordered-float", + "serde", +] + [[package]] name = "serde_core" version = "1.0.228" @@ -1802,8 +2389,19 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" dependencies = [ "cfg-if", - "cpufeatures", - "digest", + "cpufeatures 0.2.17", + "digest 0.10.7", +] + +[[package]] +name = "sha2" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "446ba717509524cb3f22f17ecc096f10f4822d76ab5c0b9822c5f9c284e825f4" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "digest 0.11.3", ] [[package]] @@ -1831,6 +2429,22 @@ dependencies = [ "libc", ] +[[package]] +name = "simd_cesu8" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94f90157bb87cddf702797c5dadfa0be7d266cdf49e22da2fcaa32eff75b2c33" +dependencies = [ + "rustc_version", + "simdutf8", +] + +[[package]] +name = "simdutf8" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" + [[package]] name = "siphasher" version = "1.0.2" @@ -1920,6 +2534,9 @@ name = "sync_wrapper" version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" +dependencies = [ + "futures-core", +] [[package]] name = "synstructure" @@ -1932,6 +2549,27 @@ dependencies = [ "syn", ] +[[package]] +name = "system-configuration" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a13f3d0daba03132c0aa9767f98351b3488edc2c100cda2d2ec2b04f3d8d3c8b" +dependencies = [ + "bitflags", + "core-foundation 0.9.4", + "system-configuration-sys", +] + +[[package]] +name = "system-configuration-sys" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e1d1b10ced5ca923a1fcb8d03e96b8d3268065d724548c0211415ff6ac6bac4" +dependencies = [ + "core-foundation-sys", + "libc", +] + [[package]] name = "thiserror" version = "2.0.18" @@ -2069,6 +2707,16 @@ dependencies = [ "whoami", ] +[[package]] +name = "tokio-rustls" +version = "0.26.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61" +dependencies = [ + "rustls", + "tokio", +] + [[package]] name = "tokio-stream" version = "0.1.18" @@ -2197,6 +2845,24 @@ dependencies = [ "tower-service", ] +[[package]] +name = "tower-http" +version = "0.6.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840" +dependencies = [ + "bitflags", + "bytes", + "futures-util", + "http", + "http-body", + "pin-project-lite", + "tower", + "tower-layer", + "tower-service", + "url", +] + [[package]] name = "tower-http" version = "0.7.0" @@ -2381,6 +3047,12 @@ dependencies = [ "tracing-serde", ] +[[package]] +name = "try-lock" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" + [[package]] name = "type-map" version = "0.5.1" @@ -2438,10 +3110,16 @@ version = "0.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea" dependencies = [ - "crypto-common", + "crypto-common 0.1.7", "subtle", ] +[[package]] +name = "untrusted" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1" + [[package]] name = "url" version = "2.5.8" @@ -2490,6 +3168,25 @@ version = "0.9.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" +[[package]] +name = "walkdir" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29790946404f91d9c5d06f9874efddea1dc06c5efe94541a7d6863108e3a5e4b" +dependencies = [ + "same-file", + "winapi-util", +] + +[[package]] +name = "want" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bfa7760aed19e106de2c7c0b581b509f2f25d3dacaf737cb82ac61bc6d760b0e" +dependencies = [ + "try-lock", +] + [[package]] name = "wasi" version = "0.11.1+wasi-snapshot-preview1" @@ -2536,6 +3233,16 @@ dependencies = [ "wasm-bindgen-shared", ] +[[package]] +name = "wasm-bindgen-futures" +version = "0.4.75" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "503b14d284f2c8dac03b819967e155ea753f573586193b2b2c95990cb5d69280" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + [[package]] name = "wasm-bindgen-macro" version = "0.2.125" @@ -2578,6 +3285,25 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "web-time" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a6580f308b1fad9207618087a65c04e7a10bc77e02c8e84e9b00dd4b12fa0bb" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "webpki-root-certs" +version = "1.0.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d46a5a140e6f7afeccd8eae97eff335163939eac8b929834875168b29b3d267" +dependencies = [ + "rustls-pki-types", +] + [[package]] name = "whoami" version = "2.1.1" @@ -2591,6 +3317,15 @@ dependencies = [ "web-sys", ] +[[package]] +name = "winapi-util" +version = "0.1.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "windows-core" version = "0.62.2" @@ -2632,6 +3367,17 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" +[[package]] +name = "windows-registry" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "02752bf7fbdcce7f2a27a742f798510f3e5ad88dbe84871e5168e2120c3d5720" +dependencies = [ + "windows-link", + "windows-result", + "windows-strings", +] + [[package]] name = "windows-result" version = "0.4.1" @@ -2829,6 +3575,12 @@ dependencies = [ "synstructure", ] +[[package]] +name = "zeroize" +version = "1.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" + [[package]] name = "zerotrie" version = "0.2.4" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 1fa40dd..9e62e0c 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -7,6 +7,7 @@ publish = false [dependencies] anyhow = "1.0.102" +async-oauth2 = "0.6.0" async-trait = "0.1.89" axum = { version = "0.8.9", features = ["json", "query"] } axum-helmet = "1.0.2" @@ -37,7 +38,9 @@ fred = { default-features = false, features = ["i-keys", "i-streams"] } +futures = "0.3.32" hex = "0.4.3" +reqwest = { version = "0.13.4", features = ["json"] } serde = { version = "1.0.228", features = ["derive"] } serde_json = "1.0.150" serde_with = { @@ -45,6 +48,7 @@ serde_with = { default-features = false, features = ["macros"] } +subtle = "2.6.1" tokio = { version = "1.52.3", default-features = false, diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index f3367dd..dfec51a 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -1,36 +1,80 @@ -use axum::{Json, extract::State, http::StatusCode, response::IntoResponse}; +use axum::{ + Extension, Json, + extract::{Path, Query, State}, + http::StatusCode, + response::{IntoResponse, Redirect}, + routing, +}; +use serde::Deserialize; use crate::{ - error::{AppError, AppResult}, + api::RoutePrefix, + error::AppResult, extractors::session::{SessionMeta, UserSession}, + services::auth::oauth::OAuthProviderEnum, state::AppState, }; pub fn routes() -> axum::Router { axum::Router::new() - .route("/login", axum::routing::post(login_handler)) - .route("/user", axum::routing::get(get_user_handler)) - .route("/logout", axum::routing::post(logout_handler)) + .route("/login/{provider}", routing::get(login_handler)) + .route("/login/{provider}/callback", routing::get(callback_handler)) + .route("/user", routing::get(get_user_handler)) + .route("/logout", routing::post(logout_handler)) +} + +fn callback_path(route_prefix: &'static str, provider: OAuthProviderEnum) -> String { + format!("{route_prefix}/login/{}/callback", provider.as_str()) } async fn login_handler( - maybe_user: Option, + Path(provider): Path, + Extension(RoutePrefix(prefix)): Extension, + State(state): State, session: tower_sessions::Session, - meta: SessionMeta, +) -> AppResult { + let oauth = state.auth_service().oauth(); + let auth_url = oauth + .authorize_url(provider, &callback_path(prefix, provider), &session) + .await?; + + Ok(Redirect::to(auth_url.as_str())) +} + +#[derive(Debug, Clone, Deserialize)] +struct OAuthCallbackQuery { + code: String, + state: String, +} + +async fn callback_handler( + Path(provider): Path, + Query(query): Query, + Extension(RoutePrefix(prefix)): Extension, State(state): State, + session: tower_sessions::Session, + meta: SessionMeta, + maybe_user: Option, ) -> AppResult { - if maybe_user.is_some() { - return Err(AppError::bad_request("already logged in")); - } + let oauth = state.auth_service().oauth(); + let token = oauth + .exchange_code( + provider, + &callback_path(prefix, provider), + &session, + &query.code, + &query.state, + ) + .await?; + let user = oauth.get_user(provider, &token, maybe_user).await?; - // TODO login handling logic - let user_id = uuid::Uuid::parse_str("6976658f-8eef-4a76-ad37-46243f463726").unwrap(); state .auth_service() - .init_session(&session, &meta, &user_id) + .init_session(&session, &meta, &user.id) .await?; - Ok(format!("Logged in as {user_id}")) + Ok(Redirect::to("/api/auth/user")) + // Ok(Redirect::to(&state.config.server.base_url)) } async fn get_user_handler( diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index f1662ca..38f55d1 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -1,3 +1,4 @@ +use axum::{Extension, Router}; use axum_plugin::AdHocPlugin; use crate::state::AppState; @@ -9,11 +10,17 @@ pub mod hello; /// Adds all API routes to the server under `/api` pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("API routes").on_setup(|router, _state| { - let api_routes = axum::Router::new() - .nest("/auth", auth::routes()) + let api_routes = Router::new() + .nest( + "/auth", + auth::routes().layer(Extension(RoutePrefix("/api/auth"))), + ) .nest("/hello", hello::routes()) .nest("/health", health::routes()); Ok(router.nest("/api", api_routes)) }) } + +#[derive(Clone)] +struct RoutePrefix(pub &'static str); diff --git a/server-new/src/config.rs b/server-new/src/config.rs index cd6e85c..8803f05 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -5,7 +5,10 @@ use axum_plugin::AdHocPlugin; use figment::providers::{Env, Format, Toml}; use serde::Deserialize; -use crate::state::AppState; +use crate::{ + services::auth::oauth::{DiscordOAuthConfig, GitHubOAuthConfig}, + state::AppState, +}; /// Parsed app configuration #[derive(Debug, Clone, Deserialize)] @@ -21,6 +24,7 @@ pub struct AppConfig { pub struct ServerConfig { pub host: IpAddr, pub port: u16, + pub base_url: String, pub log_level: String, pub request_id_header: String, pub ip_header: Option, @@ -36,6 +40,8 @@ pub struct AuthConfig { pub cookie_key: String, pub cookie_name: String, pub session_length: i64, + pub github: Option, + pub discord: Option, } #[derive(Debug, Clone, Deserialize)] diff --git a/server-new/src/db/repositories/user.rs b/server-new/src/db/repositories/user.rs index 444ab1a..a1409bc 100644 --- a/server-new/src/db/repositories/user.rs +++ b/server-new/src/db/repositories/user.rs @@ -3,7 +3,11 @@ use diesel::result::Error; use diesel_async::RunQueryDsl; use uuid::Uuid; -use crate::db::{DbConnection, models::ChatRsUser, schema::users}; +use crate::db::{ + DbConnection, + models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, + schema::users, +}; pub struct UserRepository<'a> { db: &'a mut DbConnection, @@ -14,6 +18,14 @@ impl<'a> UserRepository<'a> { UserRepository { db } } + pub async fn create(&mut self, user: NewChatRsUser<'_>) -> Result { + diesel::insert_into(users::table) + .values(user) + .returning(ChatRsUser::as_returning()) + .get_result(self.db) + .await + } + pub async fn find_by_id(&mut self, id: &Uuid) -> Result, Error> { let user = users::table .filter(users::id.eq(id)) @@ -47,16 +59,16 @@ impl<'a> UserRepository<'a> { // Ok(user) // } - // pub async fn find_by_discord_id(&mut self, id: &str) -> Result, Error> { - // let user = users::table - // .filter(users::discord_id.eq(id)) - // .select(ChatRsUser::as_select()) - // .first(self.db) - // .await - // .optional()?; + pub async fn find_by_discord_id(&mut self, id: &str) -> Result, Error> { + let user = users::table + .filter(users::discord_id.eq(id)) + .select(ChatRsUser::as_select()) + .first(self.db) + .await + .optional()?; - // Ok(user) - // } + Ok(user) + } // pub async fn find_by_oidc_id(&mut self, id: &str) -> Result, Error> { // let user = users::table @@ -80,27 +92,19 @@ impl<'a> UserRepository<'a> { // Ok(user_id) // } - // pub async fn create(&mut self, user: NewChatRsUser<'_>) -> Result { - // diesel::insert_into(users::table) - // .values(user) - // .returning(ChatRsUser::as_returning()) - // .get_result(self.db) - // .await - // } - - // pub async fn update( - // &mut self, - // user_id: &Uuid, - // data: UpdateChatRsUser<'_>, - // ) -> Result { - // let updated_id: Uuid = diesel::update(users::table.find(user_id)) - // .set(data) - // .returning(users::id) - // .get_result(self.db) - // .await?; - - // Ok(updated_id) - // } + pub async fn update( + &mut self, + user_id: &Uuid, + data: UpdateChatRsUser<'_>, + ) -> Result { + let updated_id: Uuid = diesel::update(users::table.find(user_id)) + .set(data) + .returning(users::id) + .get_result(self.db) + .await?; + + Ok(updated_id) + } // pub async fn delete(&mut self, user_id: &Uuid) -> Result { // let id: Uuid = diesel::delete(users::table.find(user_id)) diff --git a/server-new/src/error.rs b/server-new/src/error.rs index dfb23e3..9f1e5fd 100644 --- a/server-new/src/error.rs +++ b/server-new/src/error.rs @@ -53,8 +53,12 @@ impl AppError { Self::new(StatusCode::BAD_REQUEST, message) } - pub fn unauthorized() -> Self { - Self::new(StatusCode::UNAUTHORIZED, "unauthorized") + pub fn unauthorized(source: impl Into) -> Self { + Self { + status: StatusCode::UNAUTHORIZED, + message: "unauthorized".into(), + source: Some(anyhow::anyhow!(source.into())), + } } pub fn not_found(message: impl Into) -> Self { diff --git a/server-new/src/extractors/session.rs b/server-new/src/extractors/session.rs index 9888b06..cbe5258 100644 --- a/server-new/src/extractors/session.rs +++ b/server-new/src/extractors/session.rs @@ -65,7 +65,7 @@ impl FromRequestParts for UserSession { match >::from_request_parts(parts, state).await? { Some(user_session) => Ok(user_session), - None => Err(AppError::unauthorized()), + None => Err(AppError::unauthorized("no active session")), } } } diff --git a/server-new/src/lib.rs b/server-new/src/lib.rs index f89cde8..e15dd49 100644 --- a/server-new/src/lib.rs +++ b/server-new/src/lib.rs @@ -12,7 +12,11 @@ mod services; mod state; pub async fn create_app() -> anyhow::Result> { + let http_client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build()?; let app = App::new() + .store(http_client) // Add shared http client .register(config::plugin()) // Extract configuration and add to state .register(plugins::database::plugin()) // Initialize database .register(plugins::redis::plugin()) // Initialize Redis diff --git a/server-new/src/main.rs b/server-new/src/main.rs index fb2a83a..3106901 100644 --- a/server-new/src/main.rs +++ b/server-new/src/main.rs @@ -42,9 +42,7 @@ fn init_logging() -> ( tracing_subscriber::reload::Handle, tracing_appender::non_blocking::WorkerGuard, ) { - let init_log_level = std::env::var("RS_CHAT_SERVER__LOG_LEVEL") - .or_else(|_| std::env::var("RS_CHAT_LOG_LEVEL")) - .unwrap_or("info".into()); + let init_log_level = std::env::var("RS_CHAT_SERVER__LOG_LEVEL").unwrap_or("info".into()); let (writer, guard) = tracing_appender::non_blocking(std::io::stdout()); let (filter_layer, filter_handle) = tracing_subscriber::reload::Layer::new(EnvFilter::new(init_log_level)); diff --git a/server-new/src/plugins/session.rs b/server-new/src/plugins/session.rs index 1a8654d..176af99 100644 --- a/server-new/src/plugins/session.rs +++ b/server-new/src/plugins/session.rs @@ -2,11 +2,11 @@ use anyhow::{Context, bail}; use axum_plugin::AdHocPlugin; use tower_sessions::{ CachingSessionStore, Expiry, SessionManagerLayer, - cookie::{Key, SameSite, time::Duration}, + cookie::{Key, SameSite}, }; use tower_sessions_redis_store::RedisStore; -use crate::{services::SessionDbStore, state::AppState}; +use crate::{services::auth::session_store::SessionDbStore, state::AppState}; const REDIS_PREFIX: &str = "rs-chat:sess:"; @@ -22,11 +22,10 @@ pub fn plugin() -> AdHocPlugin { let redis_store = RedisStore::with_prefix(state.redis.clone(), REDIS_PREFIX.to_owned()); let db_store = SessionDbStore::new(state.db_pool.clone()); let session_store = CachingSessionStore::new(redis_store, db_store); + let session_layer = SessionManagerLayer::new(session_store) .with_name(state.config.auth.cookie_name.clone()) - .with_expiry(Expiry::OnInactivity(Duration::seconds( - state.config.auth.session_length, - ))) + .with_expiry(Expiry::OnSessionEnd) .with_private(Key::derive_from(&cookie_key)) .with_path("/") .with_secure(true) diff --git a/server-new/src/services/auth/mod.rs b/server-new/src/services/auth/mod.rs index 0fc4eaa..6a36409 100644 --- a/server-new/src/services/auth/mod.rs +++ b/server-new/src/services/auth/mod.rs @@ -1,6 +1,5 @@ -use std::fmt::Debug; - use crate::{ + config::AppConfig, db::{DbPool, DbService, models::ChatRsUser}, error::{AppError, AppResult}, extractors::session::SessionMeta, @@ -8,8 +7,8 @@ use crate::{ use tower_sessions::Session; use uuid::Uuid; -mod session_store; -pub use session_store::SessionDbStore; +pub mod oauth; +pub mod session_store; /// The field used to store the user ID in the session const USER_ID_FIELD: &str = "user_id"; @@ -17,18 +16,18 @@ const USER_ID_FIELD: &str = "user_id"; const META_FIELD: &str = "meta"; pub struct AuthService<'a> { + config: &'a AppConfig, db: &'a DbPool, -} - -impl<'a> Debug for AuthService<'a> { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("AuthService").finish() - } + http_client: &'a reqwest::Client, } impl<'a> AuthService<'a> { - pub fn new(db: &'a DbPool) -> Self { - Self { db } + pub fn new(config: &'a AppConfig, http_client: &'a reqwest::Client, db: &'a DbPool) -> Self { + Self { + config, + http_client, + db, + } } /// Initialize a new logged-in session for the given user @@ -40,6 +39,9 @@ impl<'a> AuthService<'a> { ) -> AppResult<()> { session.insert(USER_ID_FIELD, user_id).await?; session.insert(META_FIELD, meta).await?; + session.set_expiry(Some(tower_sessions::Expiry::OnInactivity( + tower_sessions::cookie::time::Duration::seconds(self.config.auth.session_length), + ))); Ok(()) } @@ -64,4 +66,13 @@ impl<'a> AuthService<'a> { session.flush().await?; Ok(()) } + + /// Access OAuth functions + pub fn oauth(self) -> oauth::OAuthService<'a> { + oauth::OAuthService { + config: self.config, + db: self.db, + http_client: self.http_client, + } + } } diff --git a/server-new/src/services/auth/oauth/discord.rs b/server-new/src/services/auth/oauth/discord.rs new file mode 100644 index 0000000..e7babbf --- /dev/null +++ b/server-new/src/services/auth/oauth/discord.rs @@ -0,0 +1,108 @@ +use futures::future::BoxFuture; +use serde::Deserialize; + +use crate::{ + db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, + error::AppResult, + services::auth::oauth::OAuthProvider, +}; + +#[derive(Clone, Debug, Deserialize)] +pub struct DiscordOAuthConfig { + client_id: u64, + client_secret: String, +} + +pub struct DiscordOAuthProvider { + config: DiscordOAuthConfig, +} + +impl DiscordOAuthProvider { + pub fn new(config: &DiscordOAuthConfig) -> Self { + Self { + config: config.clone(), + } + } +} + +impl OAuthProvider for DiscordOAuthProvider { + fn get_authorize_url(&self) -> &str { + "https://discord.com/oauth2/authorize" + } + + fn get_token_url(&self) -> &str { + "https://discord.com/api/oauth2/token" + } + + fn get_scopes(&self) -> Vec<&str> { + vec!["identify"] + } + + fn get_user_info_url(&self) -> &str { + "https://discord.com/api/v9/users/@me" + } + + fn get_client_id(&self) -> String { + self.config.client_id.to_string() + } + + fn get_client_secret(&self) -> String { + self.config.client_secret.clone() + } + + fn extract_user_data(&self, user_info: serde_json::Value) -> anyhow::Result { + let user_info: DiscordUserInfo = serde_json::from_value(user_info)?; + let avatar_url = user_info.avatar.as_ref().map(|avatar| { + format!( + "https://cdn.discordapp.com/avatars/{}/{}.png", + user_info.id, avatar + ) + }); + + Ok(super::UserData { + id: user_info.id, + name: user_info.global_name.unwrap_or_else(|| user_info.username), + avatar_url, + }) + } + + fn find_linked_user<'a>( + &self, + db: &'a mut crate::db::DbService, + user_data: &'a super::UserData, + ) -> BoxFuture<'a, AppResult>> { + Box::pin(async move { + let user = db.users().find_by_discord_id(&user_data.id).await?; + Ok(user) + }) + } + + fn is_user_linked(&self, user: &ChatRsUser) -> bool { + user.discord_id.is_some() + } + + fn create_update_user<'a>(&self, user_data: &'a super::UserData) -> UpdateChatRsUser<'a> { + UpdateChatRsUser { + discord_id: Some(&user_data.id), + ..Default::default() + } + } + + fn create_new_user<'a>(&self, user_data: &'a super::UserData) -> NewChatRsUser<'a> { + NewChatRsUser { + discord_id: Some(&user_data.id), + name: &user_data.name, + avatar_url: user_data.avatar_url.as_deref(), + ..Default::default() + } + } +} + +/// User info returned from Discord API +#[derive(Debug, Deserialize)] +pub struct DiscordUserInfo { + id: String, + username: String, + global_name: Option, + avatar: Option, +} diff --git a/server-new/src/services/auth/oauth/github.rs b/server-new/src/services/auth/oauth/github.rs new file mode 100644 index 0000000..65bc533 --- /dev/null +++ b/server-new/src/services/auth/oauth/github.rs @@ -0,0 +1,109 @@ +use futures::future::BoxFuture; +use serde::Deserialize; + +use crate::{ + db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, + error::AppResult, + services::auth::oauth::OAuthProvider, +}; + +#[derive(Clone, Debug, Deserialize)] +pub struct GitHubOAuthConfig { + client_id: String, + client_secret: String, +} + +pub struct GitHubOAuthProvider { + config: GitHubOAuthConfig, +} + +impl GitHubOAuthProvider { + pub fn new(config: &GitHubOAuthConfig) -> Self { + Self { + config: config.clone(), + } + } +} + +impl OAuthProvider for GitHubOAuthProvider { + fn get_authorize_url(&self) -> &str { + "https://github.com/login/oauth/authorize" + } + + fn get_token_url(&self) -> &str { + "https://github.com/login/oauth/access_token" + } + + fn get_scopes(&self) -> Vec<&str> { + vec!["user:read"] + } + + fn get_user_info_url(&self) -> &str { + "https://api.github.com/user" + } + + fn get_client_id(&self) -> String { + self.config.client_id.to_owned() + } + + fn get_client_secret(&self) -> String { + self.config.client_secret.to_owned() + } + + fn create_request_headers(&self) -> Vec<(&'static str, &'static str)> { + vec![ + ("Accept", "application/vnd.github+json"), + ("User-Agent", "fa-sharp/rs-chat"), + ] + } + + fn extract_user_data(&self, user_info: serde_json::Value) -> anyhow::Result { + let info: GitHubUserInfo = serde_json::from_value(user_info)?; + + Ok(super::UserData { + id: info.id.to_string(), + name: info.name.unwrap_or(info.login), + avatar_url: info.avatar_url, + }) + } + + fn find_linked_user<'a>( + &self, + db: &'a mut crate::db::DbService, + user_data: &'a super::UserData, + ) -> BoxFuture<'a, AppResult>> { + Box::pin(async move { + let user = db.users().find_by_github_id(&user_data.id).await?; + Ok(user) + }) + } + + fn is_user_linked(&self, user: &ChatRsUser) -> bool { + user.github_id.is_some() + } + + fn create_update_user<'a>(&self, user_data: &'a super::UserData) -> UpdateChatRsUser<'a> { + UpdateChatRsUser { + github_id: Some(&user_data.id), + ..Default::default() + } + } + + fn create_new_user<'a>(&self, user_data: &'a super::UserData) -> NewChatRsUser<'a> { + NewChatRsUser { + github_id: Some(&user_data.id), + name: &user_data.name, + avatar_url: user_data.avatar_url.as_deref(), + ..Default::default() + } + } +} + +/// User info returned from GitHub API +#[derive(Debug, Deserialize)] +struct GitHubUserInfo { + id: u64, + login: String, + name: Option, + avatar_url: Option, +} diff --git a/server-new/src/services/auth/oauth/mod.rs b/server-new/src/services/auth/oauth/mod.rs new file mode 100644 index 0000000..8eaefaa --- /dev/null +++ b/server-new/src/services/auth/oauth/mod.rs @@ -0,0 +1,228 @@ +use anyhow::Context; +use futures::future::BoxFuture; +use oauth2::{StandardToken, Token}; +use serde::{Deserialize, Serialize}; +use subtle::ConstantTimeEq; +use tower_sessions::Session; + +use crate::{ + config::AppConfig, + db::{ + DbPool, DbService, + models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, + }, + error::{AppError, AppResult}, + extractors::session::UserSession, +}; + +mod discord; +mod github; + +pub use discord::DiscordOAuthConfig; +pub use github::GitHubOAuthConfig; + +/// Supported OAuth providers +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum OAuthProviderEnum { + Github, + Discord, +} +impl OAuthProviderEnum { + pub fn as_str(&self) -> &str { + match self { + OAuthProviderEnum::Github => "github", + OAuthProviderEnum::Discord => "discord", + } + } +} + +/// Trait for all OAuth providers +pub trait OAuthProvider: Send + Sync { + fn get_scopes(&self) -> Vec<&str>; + fn get_authorize_url(&self) -> &str; + fn get_token_url(&self) -> &str; + fn get_user_info_url(&self) -> &str; + fn get_client_id(&self) -> String; + fn get_client_secret(&self) -> String; + fn create_request_headers(&self) -> Vec<(&'static str, &'static str)> { + vec![] + } + fn extract_user_data(&self, user_info: serde_json::Value) -> anyhow::Result; + fn find_linked_user<'a>( + &self, + db: &'a mut DbService, + user_data: &'a UserData, + ) -> BoxFuture<'a, AppResult>>; + fn is_user_linked(&self, user: &ChatRsUser) -> bool; + fn create_update_user<'a>(&self, user_data: &'a UserData) -> UpdateChatRsUser<'a>; + fn create_new_user<'a>(&self, user_data: &'a UserData) -> NewChatRsUser<'a>; +} + +/// OAuth functions +pub struct OAuthService<'a> { + pub(super) config: &'a AppConfig, + pub(super) db: &'a DbPool, + pub(super) http_client: &'a reqwest::Client, +} + +impl OAuthService<'_> { + const SESS_STATE_FIELD: &'static str = "oauth_state"; + const SESS_PKCE_FIELD: &'static str = "oauth_verifier"; + + pub async fn authorize_url( + &self, + provider: OAuthProviderEnum, + callback_path: &str, + session: &Session, + ) -> AppResult { + let provider = self.get_provider(provider)?; + let client = self.get_oauth_client(provider.as_ref(), callback_path)?; + + let state = oauth2::State::new_random(); + let pkce_verifier = oauth2::PkceCodeVerifierS256::new_random(); + let mut auth_url = client.authorize_url(&state); + auth_url + .query_pairs_mut() + .extend_pairs(pkce_verifier.authorize_url_params()); + + session.insert(Self::SESS_STATE_FIELD, state).await?; + session.insert(Self::SESS_PKCE_FIELD, pkce_verifier).await?; + session.set_expiry(Some(tower_sessions::Expiry::OnSessionEnd)); + + Ok(auth_url) + } + + pub async fn exchange_code( + &self, + provider: OAuthProviderEnum, + callback_path: &str, + session: &Session, + code: &str, + returned_state: &str, + ) -> AppResult { + // Get saved state and code verifier from session + let saved_state = session + .remove::(Self::SESS_STATE_FIELD) + .await? + .ok_or_else(|| AppError::unauthorized("no state in session"))?; + let pkce_verifier = session + .remove::(Self::SESS_PKCE_FIELD) + .await? + .ok_or_else(|| AppError::unauthorized("no PKCE verifier in session"))?; + + // Verify state + if saved_state.ct_ne(returned_state.as_bytes()).into() { + return Err(AppError::unauthorized("state parameter doesn't match")); + } + + // Exchange code for token + let provider = self.get_provider(provider)?; + let client = self.get_oauth_client(provider.as_ref(), callback_path)?; + let response = client + .exchange_code(code) + .param("code_verifier", String::from(pkce_verifier)) + .with_reqwest_client(&self.http_client) + .execute::() + .await + .map_err(|err| AppError::unauthorized(format!("token exchange failed: {err}")))?; + + Ok(response) + } + + pub async fn get_user( + &self, + provider: OAuthProviderEnum, + token: &oauth2::StandardToken, + active_session: Option, + ) -> AppResult { + let provider = self.get_provider(provider)?; + let mut user_info_request = self + .http_client + .get(provider.get_user_info_url()) + .bearer_auth(token.access_token().as_ref()); + for (name, value) in provider.create_request_headers() { + user_info_request = user_info_request.header(name, value); + } + let user_info_response = user_info_request.send().await.context("request failed")?; + if !user_info_response.status().is_success() { + let error = + anyhow::anyhow!("failed to get user: {:?}", user_info_response.text().await); + return Err(AppError::internal(error)); + } + let user_data = provider + .extract_user_data(user_info_response.json().await.context("request failed")?) + .context("unable to extract user data from response")?; + + let mut db = DbService::from_pool(self.db).await?; + let user = match provider.find_linked_user(&mut db, &user_data).await? { + Some(existing_user) => existing_user, + None => match active_session { + None => { + let new_user = provider.create_new_user(&user_data); + db.users().create(new_user).await? + } + Some(sess) => match db.users().find_by_id(&sess.user_id).await? { + Some(user) if provider.is_user_linked(&user) => { + return Err(AppError::bad_request("user already linked to provider")); + } + Some(user) => { + // Link logged-in user to new provider + let update_user = provider.create_update_user(&user_data); + db.users().update(&user.id, update_user).await?; + user + } + None => { + return Err(AppError::internal(anyhow::anyhow!("user not found"))); + } + }, + }, + }; + + Ok(user) + } + + fn get_provider(&self, provider: OAuthProviderEnum) -> AppResult> { + let provider: Option> = match provider { + OAuthProviderEnum::Github => match self.config.auth.github { + Some(ref c) => Some(Box::new(github::GitHubOAuthProvider::new(c))), + None => None, + }, + OAuthProviderEnum::Discord => match self.config.auth.discord { + Some(ref c) => Some(Box::new(discord::DiscordOAuthProvider::new(c))), + None => None, + }, + }; + + provider.ok_or_else(|| AppError::bad_request("unsupported OAuth provider")) + } + + fn get_oauth_client( + &self, + provider: &dyn OAuthProvider, + callback_path: &str, + ) -> anyhow::Result { + let mut client = oauth2::Client::new( + provider.get_client_id(), + provider.get_authorize_url().parse()?, + provider.get_token_url().parse()?, + ); + client.set_client_secret(provider.get_client_secret()); + client.set_redirect_url( + format!("{}{}", &self.config.server.base_url, callback_path).parse()?, + ); + for scope in provider.get_scopes() { + client.add_scope(scope); + } + + Ok(client) + } +} + +/// Common OAuth user data structure +#[derive(Debug)] +pub struct UserData { + pub id: String, + pub name: String, + pub avatar_url: Option, +} diff --git a/server-new/src/services/mod.rs b/server-new/src/services/mod.rs index 07f4c79..0e4a05d 100644 --- a/server-new/src/services/mod.rs +++ b/server-new/src/services/mod.rs @@ -1,3 +1 @@ -mod auth; - -pub use auth::{AuthService, SessionDbStore}; +pub mod auth; diff --git a/server-new/src/state.rs b/server-new/src/state.rs index fd7b37d..ede4370 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -4,7 +4,7 @@ use std::{ops::Deref, sync::Arc}; use axum_plugin::{AppState, TypeMap}; -use crate::{config::AppConfig, db::DbPool, services::AuthService}; +use crate::{config::AppConfig, db::DbPool, services::auth::AuthService}; /// App state stored in the Axum router #[derive(Clone)] @@ -13,13 +13,14 @@ pub struct AppState(Arc); #[derive(AppState)] pub struct AppStateInner { pub config: AppConfig, + pub http_client: reqwest::Client, pub db_pool: DbPool, pub redis: fred::prelude::Pool, } impl AppState { pub fn auth_service(&self) -> AuthService<'_> { - AuthService::new(&self.db_pool) + AuthService::new(&self.config, &self.http_client, &self.db_pool) } } From 47382f946b62d6ecee25cad14271d19c38da9f22 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 18:12:35 -0400 Subject: [PATCH 032/111] oauth tweaks --- server-new/src/db/repositories/session.rs | 3 +- server-new/src/services/auth/mod.rs | 7 +- server-new/src/services/auth/oauth/mod.rs | 172 +--------------- server-new/src/services/auth/oauth/service.rs | 185 ++++++++++++++++++ server-new/src/services/auth/session_store.rs | 34 ++-- 5 files changed, 211 insertions(+), 190 deletions(-) create mode 100644 server-new/src/services/auth/oauth/service.rs diff --git a/server-new/src/db/repositories/session.rs b/server-new/src/db/repositories/session.rs index ecfad57..6e14206 100644 --- a/server-new/src/db/repositories/session.rs +++ b/server-new/src/db/repositories/session.rs @@ -44,8 +44,7 @@ impl<'a> SessionRepository<'a> { data: &HashMap, expires_at: UtcDateTime, ) -> QueryResult { - diesel::update(auth_sessions::table) - .filter(auth_sessions::id.eq(session_id)) + diesel::update(auth_sessions::table.find(session_id)) .set(UpdateChatRsAuthSession { data: AuthSessionData(data.to_owned()), expires_at, diff --git a/server-new/src/services/auth/mod.rs b/server-new/src/services/auth/mod.rs index 6a36409..4f33343 100644 --- a/server-new/src/services/auth/mod.rs +++ b/server-new/src/services/auth/mod.rs @@ -37,6 +37,7 @@ impl<'a> AuthService<'a> { meta: &SessionMeta, user_id: &Uuid, ) -> AppResult<()> { + session.cycle_id().await?; session.insert(USER_ID_FIELD, user_id).await?; session.insert(META_FIELD, meta).await?; session.set_expiry(Some(tower_sessions::Expiry::OnInactivity( @@ -69,10 +70,6 @@ impl<'a> AuthService<'a> { /// Access OAuth functions pub fn oauth(self) -> oauth::OAuthService<'a> { - oauth::OAuthService { - config: self.config, - db: self.db, - http_client: self.http_client, - } + oauth::OAuthService::new(self.config, self.db, self.http_client) } } diff --git a/server-new/src/services/auth/oauth/mod.rs b/server-new/src/services/auth/oauth/mod.rs index 8eaefaa..4d33934 100644 --- a/server-new/src/services/auth/oauth/mod.rs +++ b/server-new/src/services/auth/oauth/mod.rs @@ -1,25 +1,21 @@ -use anyhow::Context; use futures::future::BoxFuture; -use oauth2::{StandardToken, Token}; use serde::{Deserialize, Serialize}; -use subtle::ConstantTimeEq; -use tower_sessions::Session; use crate::{ - config::AppConfig, db::{ - DbPool, DbService, + DbService, models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, }, - error::{AppError, AppResult}, - extractors::session::UserSession, + error::AppResult, }; mod discord; mod github; +mod service; pub use discord::DiscordOAuthConfig; pub use github::GitHubOAuthConfig; +pub use service::OAuthService; /// Supported OAuth providers #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] @@ -59,166 +55,6 @@ pub trait OAuthProvider: Send + Sync { fn create_new_user<'a>(&self, user_data: &'a UserData) -> NewChatRsUser<'a>; } -/// OAuth functions -pub struct OAuthService<'a> { - pub(super) config: &'a AppConfig, - pub(super) db: &'a DbPool, - pub(super) http_client: &'a reqwest::Client, -} - -impl OAuthService<'_> { - const SESS_STATE_FIELD: &'static str = "oauth_state"; - const SESS_PKCE_FIELD: &'static str = "oauth_verifier"; - - pub async fn authorize_url( - &self, - provider: OAuthProviderEnum, - callback_path: &str, - session: &Session, - ) -> AppResult { - let provider = self.get_provider(provider)?; - let client = self.get_oauth_client(provider.as_ref(), callback_path)?; - - let state = oauth2::State::new_random(); - let pkce_verifier = oauth2::PkceCodeVerifierS256::new_random(); - let mut auth_url = client.authorize_url(&state); - auth_url - .query_pairs_mut() - .extend_pairs(pkce_verifier.authorize_url_params()); - - session.insert(Self::SESS_STATE_FIELD, state).await?; - session.insert(Self::SESS_PKCE_FIELD, pkce_verifier).await?; - session.set_expiry(Some(tower_sessions::Expiry::OnSessionEnd)); - - Ok(auth_url) - } - - pub async fn exchange_code( - &self, - provider: OAuthProviderEnum, - callback_path: &str, - session: &Session, - code: &str, - returned_state: &str, - ) -> AppResult { - // Get saved state and code verifier from session - let saved_state = session - .remove::(Self::SESS_STATE_FIELD) - .await? - .ok_or_else(|| AppError::unauthorized("no state in session"))?; - let pkce_verifier = session - .remove::(Self::SESS_PKCE_FIELD) - .await? - .ok_or_else(|| AppError::unauthorized("no PKCE verifier in session"))?; - - // Verify state - if saved_state.ct_ne(returned_state.as_bytes()).into() { - return Err(AppError::unauthorized("state parameter doesn't match")); - } - - // Exchange code for token - let provider = self.get_provider(provider)?; - let client = self.get_oauth_client(provider.as_ref(), callback_path)?; - let response = client - .exchange_code(code) - .param("code_verifier", String::from(pkce_verifier)) - .with_reqwest_client(&self.http_client) - .execute::() - .await - .map_err(|err| AppError::unauthorized(format!("token exchange failed: {err}")))?; - - Ok(response) - } - - pub async fn get_user( - &self, - provider: OAuthProviderEnum, - token: &oauth2::StandardToken, - active_session: Option, - ) -> AppResult { - let provider = self.get_provider(provider)?; - let mut user_info_request = self - .http_client - .get(provider.get_user_info_url()) - .bearer_auth(token.access_token().as_ref()); - for (name, value) in provider.create_request_headers() { - user_info_request = user_info_request.header(name, value); - } - let user_info_response = user_info_request.send().await.context("request failed")?; - if !user_info_response.status().is_success() { - let error = - anyhow::anyhow!("failed to get user: {:?}", user_info_response.text().await); - return Err(AppError::internal(error)); - } - let user_data = provider - .extract_user_data(user_info_response.json().await.context("request failed")?) - .context("unable to extract user data from response")?; - - let mut db = DbService::from_pool(self.db).await?; - let user = match provider.find_linked_user(&mut db, &user_data).await? { - Some(existing_user) => existing_user, - None => match active_session { - None => { - let new_user = provider.create_new_user(&user_data); - db.users().create(new_user).await? - } - Some(sess) => match db.users().find_by_id(&sess.user_id).await? { - Some(user) if provider.is_user_linked(&user) => { - return Err(AppError::bad_request("user already linked to provider")); - } - Some(user) => { - // Link logged-in user to new provider - let update_user = provider.create_update_user(&user_data); - db.users().update(&user.id, update_user).await?; - user - } - None => { - return Err(AppError::internal(anyhow::anyhow!("user not found"))); - } - }, - }, - }; - - Ok(user) - } - - fn get_provider(&self, provider: OAuthProviderEnum) -> AppResult> { - let provider: Option> = match provider { - OAuthProviderEnum::Github => match self.config.auth.github { - Some(ref c) => Some(Box::new(github::GitHubOAuthProvider::new(c))), - None => None, - }, - OAuthProviderEnum::Discord => match self.config.auth.discord { - Some(ref c) => Some(Box::new(discord::DiscordOAuthProvider::new(c))), - None => None, - }, - }; - - provider.ok_or_else(|| AppError::bad_request("unsupported OAuth provider")) - } - - fn get_oauth_client( - &self, - provider: &dyn OAuthProvider, - callback_path: &str, - ) -> anyhow::Result { - let mut client = oauth2::Client::new( - provider.get_client_id(), - provider.get_authorize_url().parse()?, - provider.get_token_url().parse()?, - ); - client.set_client_secret(provider.get_client_secret()); - client.set_redirect_url( - format!("{}{}", &self.config.server.base_url, callback_path).parse()?, - ); - for scope in provider.get_scopes() { - client.add_scope(scope); - } - - Ok(client) - } -} - /// Common OAuth user data structure #[derive(Debug)] pub struct UserData { diff --git a/server-new/src/services/auth/oauth/service.rs b/server-new/src/services/auth/oauth/service.rs new file mode 100644 index 0000000..e5614d6 --- /dev/null +++ b/server-new/src/services/auth/oauth/service.rs @@ -0,0 +1,185 @@ +use anyhow::Context; +use oauth2::{StandardToken, Token}; +use subtle::ConstantTimeEq; +use tower_sessions::Session; + +use crate::{ + config::AppConfig, + db::{DbPool, DbService, models::ChatRsUser}, + error::{AppError, AppResult}, + extractors::session::UserSession, + services::auth::oauth::{OAuthProvider, OAuthProviderEnum}, +}; + +/// OAuth functions +pub struct OAuthService<'a> { + config: &'a AppConfig, + db: &'a DbPool, + http_client: &'a reqwest::Client, +} + +impl<'a> OAuthService<'a> { + const SESS_STATE_FIELD: &'static str = "oauth_state"; + const SESS_PKCE_FIELD: &'static str = "oauth_verifier"; + + pub fn new(config: &'a AppConfig, db: &'a DbPool, http_client: &'a reqwest::Client) -> Self { + Self { + config, + db, + http_client, + } + } + + pub async fn authorize_url( + &self, + provider: OAuthProviderEnum, + callback_path: &str, + session: &Session, + ) -> AppResult { + let provider = self.get_provider(provider)?; + let client = self.get_oauth_client(provider.as_ref(), callback_path)?; + + let state = oauth2::State::new_random(); + let pkce_verifier = oauth2::PkceCodeVerifierS256::new_random(); + let mut auth_url = client.authorize_url(&state); + auth_url + .query_pairs_mut() + .extend_pairs(pkce_verifier.authorize_url_params()); + + session.insert(Self::SESS_STATE_FIELD, state).await?; + session.insert(Self::SESS_PKCE_FIELD, pkce_verifier).await?; + + Ok(auth_url) + } + + pub async fn exchange_code( + &self, + provider: OAuthProviderEnum, + callback_path: &str, + session: &Session, + code: &str, + returned_state: &str, + ) -> AppResult { + // Get saved state and code verifier from session + let saved_state = session + .remove::(Self::SESS_STATE_FIELD) + .await? + .ok_or_else(|| AppError::unauthorized("no state in session"))?; + let pkce_verifier = session + .remove::(Self::SESS_PKCE_FIELD) + .await? + .ok_or_else(|| AppError::unauthorized("no PKCE verifier in session"))?; + + // Verify state + if saved_state.ct_ne(returned_state.as_bytes()).into() { + return Err(AppError::unauthorized("state parameter doesn't match")); + } + + // Exchange code for token + let provider = self.get_provider(provider)?; + let client = self.get_oauth_client(provider.as_ref(), callback_path)?; + let response = client + .exchange_code(code) + .param("code_verifier", String::from(pkce_verifier)) + .with_reqwest_client(&self.http_client) + .execute::() + .await + .map_err(|err| AppError::unauthorized(format!("token exchange failed: {err}")))?; + + Ok(response) + } + + pub async fn get_user( + &self, + provider: OAuthProviderEnum, + token: &oauth2::StandardToken, + active_session: Option, + ) -> AppResult { + let provider = self.get_provider(provider)?; + let mut user_info_request = self + .http_client + .get(provider.get_user_info_url()) + .bearer_auth(token.access_token().as_ref()); + for (name, value) in provider.create_request_headers() { + user_info_request = user_info_request.header(name, value); + } + let user_info_response = user_info_request.send().await.context("request failed")?; + if !user_info_response.status().is_success() { + let error = + anyhow::anyhow!("failed to get user: {:?}", user_info_response.text().await); + return Err(AppError::internal(error)); + } + let user_data = provider + .extract_user_data(user_info_response.json().await.context("request failed")?) + .context("unable to extract user data from response")?; + + let mut db = DbService::from_pool(self.db).await?; + let user = match provider.find_linked_user(&mut db, &user_data).await? { + Some(existing_user) => { + if active_session.is_some_and(|sess| sess.user_id != existing_user.id) { + return Err(AppError::unauthorized("cannot switch users via OAuth")); + } else { + existing_user + } + } + None => match active_session { + None => { + let new_user = provider.create_new_user(&user_data); + db.users().create(new_user).await? + } + Some(sess) => match db.users().find_by_id(&sess.user_id).await? { + Some(user) if provider.is_user_linked(&user) => { + return Err(AppError::bad_request("user already linked to provider")); + } + Some(user) => { + // Link logged-in user to new provider + let update_user = provider.create_update_user(&user_data); + db.users().update(&user.id, update_user).await?; + user + } + None => { + return Err(AppError::internal(anyhow::anyhow!("user not found"))); + } + }, + }, + }; + + Ok(user) + } + + fn get_provider(&self, provider: OAuthProviderEnum) -> AppResult> { + let provider: Option> = match provider { + OAuthProviderEnum::Github => match self.config.auth.github { + Some(ref c) => Some(Box::new(super::github::GitHubOAuthProvider::new(c))), + None => None, + }, + OAuthProviderEnum::Discord => match self.config.auth.discord { + Some(ref c) => Some(Box::new(super::discord::DiscordOAuthProvider::new(c))), + None => None, + }, + }; + + provider.ok_or_else(|| AppError::bad_request("unsupported OAuth provider")) + } + + fn get_oauth_client( + &self, + provider: &dyn OAuthProvider, + callback_path: &str, + ) -> anyhow::Result { + let mut client = oauth2::Client::new( + provider.get_client_id(), + provider.get_authorize_url().parse()?, + provider.get_token_url().parse()?, + ); + client.set_client_secret(provider.get_client_secret()); + client.set_redirect_url( + format!("{}{}", &self.config.server.base_url, callback_path).parse()?, + ); + for scope in provider.get_scopes() { + client.add_scope(scope); + } + + Ok(client) + } +} diff --git a/server-new/src/services/auth/session_store.rs b/server-new/src/services/auth/session_store.rs index caa4daa..31812e5 100644 --- a/server-new/src/services/auth/session_store.rs +++ b/server-new/src/services/auth/session_store.rs @@ -19,11 +19,20 @@ impl SessionDbStore { Self { db } } - fn get_session_uuid(&self, id: &Id) -> Uuid { + fn get_session_uuid(id: &Id) -> Uuid { Uuid::from_bytes(id.0.to_be_bytes()) } - fn convert_expiry(&self, time: OffsetDateTime) -> Result { + fn extract_user_id(record: &Record) -> Result> { + record + .data + .get(super::USER_ID_FIELD) + .map(|val| serde_json::from_value::(val.clone())) + .transpose() + .map_err(|_| Error::Encode("invalid user id field".to_owned())) + } + + fn convert_expiry(time: OffsetDateTime) -> Result { UtcDateTime::from_timestamp_secs(time.unix_timestamp()) .ok_or_else(|| Error::Backend(format!("Invalid expiry: {time}"))) } @@ -45,14 +54,9 @@ impl std::fmt::Debug for SessionDbStore { impl SessionStore for SessionDbStore { /// Creates a new session in the store with the provided session record. async fn create(&self, record: &mut Record) -> Result<()> { - let session_id = self.get_session_uuid(&record.id); - let user_id = record - .data - .get(super::USER_ID_FIELD) - .map(|val| serde_json::from_value::(val.clone())) - .transpose() - .map_err(|_| Error::Encode("invalid user id field".to_owned()))?; - let expires_at = self.convert_expiry(record.expiry_date)?; + let session_id = Self::get_session_uuid(&record.id); + let user_id = Self::extract_user_id(record)?; + let expires_at = Self::convert_expiry(record.expiry_date)?; let mut db = self.get_db().await?; db.sessions() @@ -67,10 +71,10 @@ impl SessionStore for SessionDbStore { /// /// This method is intended for updating the state of an existing session. async fn save(&self, record: &Record) -> Result<()> { - let session_id = self.get_session_uuid(&record.id); - let mut db = self.get_db().await?; - let expires_at = self.convert_expiry(record.expiry_date)?; + let session_id = Self::get_session_uuid(&record.id); + let expires_at = Self::convert_expiry(record.expiry_date)?; + let mut db = self.get_db().await?; db.sessions() .update(&session_id, &record.data, expires_at) .await @@ -85,7 +89,7 @@ impl SessionStore for SessionDbStore { /// does not exist or has been invalidated (e.g., expired), `None` is /// returned. async fn load(&self, session_id: &Id) -> Result> { - let session_id = self.get_session_uuid(&session_id); + let session_id = Self::get_session_uuid(&session_id); let mut db = self.get_db().await?; match db.sessions().find_active_by_id(&session_id).await { @@ -104,7 +108,7 @@ impl SessionStore for SessionDbStore { /// /// If the session exists, it is removed from the store. async fn delete(&self, session_id: &Id) -> Result<()> { - let session_id = self.get_session_uuid(session_id); + let session_id = Self::get_session_uuid(session_id); let mut db = self.get_db().await?; db.sessions() From 3ef6526df474fbf53cd3f84c0572f18e180f9e1b Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 18:42:26 -0400 Subject: [PATCH 033/111] oauth refactor --- server-new/src/api/auth.rs | 9 +- server-new/src/extractors/session.rs | 30 +-- server-new/src/services/auth/mod.rs | 38 +--- server-new/src/services/auth/oauth/mod.rs | 190 +++++++++++++++++- server-new/src/services/auth/oauth/service.rs | 185 ----------------- server-new/src/services/auth/session.rs | 67 ++++++ server-new/src/services/auth/session_store.rs | 11 +- server-new/src/services/auth/types.rs | 26 +++ 8 files changed, 295 insertions(+), 261 deletions(-) delete mode 100644 server-new/src/services/auth/oauth/service.rs create mode 100644 server-new/src/services/auth/session.rs create mode 100644 server-new/src/services/auth/types.rs diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index dfec51a..5ca3734 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -10,8 +10,7 @@ use serde::Deserialize; use crate::{ api::RoutePrefix, error::AppResult, - extractors::session::{SessionMeta, UserSession}, - services::auth::oauth::OAuthProviderEnum, + services::auth::{SessionMeta, UserSession, oauth::OAuthProviderEnum}, state::AppState, }; @@ -67,10 +66,10 @@ async fn callback_handler( ) .await?; let user = oauth.get_user(provider, &token, maybe_user).await?; - state .auth_service() - .init_session(&session, &meta, &user.id) + .session() + .login(&session, &meta, &user.id) .await?; Ok(Redirect::to("/api/auth/user")) @@ -89,6 +88,6 @@ async fn logout_handler( session: tower_sessions::Session, State(state): State, ) -> AppResult { - state.auth_service().logout_user(&session).await?; + state.auth_service().session().logout(&session).await?; Ok(StatusCode::NO_CONTENT) } diff --git a/server-new/src/extractors/session.rs b/server-new/src/extractors/session.rs index cbe5258..ac9ef99 100644 --- a/server-new/src/extractors/session.rs +++ b/server-new/src/extractors/session.rs @@ -7,11 +7,13 @@ use axum::{ extract::{ConnectInfo, FromRequestParts, OptionalFromRequestParts}, http::header, }; -use serde::{Deserialize, Serialize}; use tower_sessions::Session; -use uuid::Uuid; -use crate::{db::UtcDateTime, error::AppError, state::AppState}; +use crate::{ + error::AppError, + services::auth::{SessionMeta, UserSession}, + state::AppState, +}; /// Active user session data. /// @@ -20,25 +22,6 @@ use crate::{db::UtcDateTime, error::AppError, state::AppState}; /// and `None` otherwise. /// - If used as `UserSession`, request will automatically return an unauthorized error /// if there is no active session. -#[derive(Debug, Serialize, Deserialize)] -pub struct UserSession { - pub user_id: Uuid, -} - -/// Session metadata extracted on login. -#[derive(Debug, Serialize, Deserialize)] -pub struct SessionMeta { - pub start_time: UtcDateTime, - pub ip: Option, - pub user_agent: Option, -} - -impl UserSession { - pub fn new(user_id: Uuid) -> Self { - Self { user_id } - } -} - impl OptionalFromRequestParts for UserSession { type Rejection = AppError; @@ -49,9 +32,8 @@ impl OptionalFromRequestParts for UserSession { let session = Session::from_request_parts(parts, state) .await .map_err(|(_, msg)| AppError::internal(anyhow::anyhow!(msg)))?; - let user_id = state.auth_service().extract_user_id(session).await?; - Ok(user_id.map(UserSession::new)) + state.auth_service().session().user_session(&session).await } } diff --git a/server-new/src/services/auth/mod.rs b/server-new/src/services/auth/mod.rs index 4f33343..8b0ffb3 100644 --- a/server-new/src/services/auth/mod.rs +++ b/server-new/src/services/auth/mod.rs @@ -2,18 +2,15 @@ use crate::{ config::AppConfig, db::{DbPool, DbService, models::ChatRsUser}, error::{AppError, AppResult}, - extractors::session::SessionMeta, }; -use tower_sessions::Session; use uuid::Uuid; pub mod oauth; +pub mod session; pub mod session_store; +mod types; -/// The field used to store the user ID in the session -const USER_ID_FIELD: &str = "user_id"; -/// The field used to store the user session metadata -const META_FIELD: &str = "meta"; +pub use types::*; pub struct AuthService<'a> { config: &'a AppConfig, @@ -30,28 +27,6 @@ impl<'a> AuthService<'a> { } } - /// Initialize a new logged-in session for the given user - pub async fn init_session( - &self, - session: &Session, - meta: &SessionMeta, - user_id: &Uuid, - ) -> AppResult<()> { - session.cycle_id().await?; - session.insert(USER_ID_FIELD, user_id).await?; - session.insert(META_FIELD, meta).await?; - session.set_expiry(Some(tower_sessions::Expiry::OnInactivity( - tower_sessions::cookie::time::Duration::seconds(self.config.auth.session_length), - ))); - - Ok(()) - } - - /// Extract the current user ID if this is an active user session - pub async fn extract_user_id(&self, session: Session) -> AppResult> { - Ok(session.get::(USER_ID_FIELD).await?) - } - /// Get the user from the database with the given ID, or return /// an internal error if not found pub async fn get_user(&self, id: &Uuid) -> AppResult { @@ -62,10 +37,9 @@ impl<'a> AuthService<'a> { } } - /// Logout the user, deleting the current session - pub async fn logout_user(&self, session: &Session) -> AppResult<()> { - session.flush().await?; - Ok(()) + /// Access session functions. + pub fn session(&self) -> session::AuthSessionService { + session::AuthSessionService::new(self.config.auth.session_length) } /// Access OAuth functions diff --git a/server-new/src/services/auth/oauth/mod.rs b/server-new/src/services/auth/oauth/mod.rs index 4d33934..95427c5 100644 --- a/server-new/src/services/auth/oauth/mod.rs +++ b/server-new/src/services/auth/oauth/mod.rs @@ -1,21 +1,24 @@ +use anyhow::Context; use futures::future::BoxFuture; +use oauth2::{StandardToken, Token}; use serde::{Deserialize, Serialize}; +use subtle::ConstantTimeEq; +use tower_sessions::Session; use crate::{ + config::AppConfig, db::{ - DbService, + DbPool, DbService, models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, }, - error::AppResult, + error::{AppError, AppResult}, + services::auth::UserSession, }; - mod discord; mod github; -mod service; pub use discord::DiscordOAuthConfig; pub use github::GitHubOAuthConfig; -pub use service::OAuthService; /// Supported OAuth providers #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] @@ -62,3 +65,180 @@ pub struct UserData { pub name: String, pub avatar_url: Option, } + +/// OAuth functions +pub struct OAuthService<'a> { + config: &'a AppConfig, + db: &'a DbPool, + http_client: &'a reqwest::Client, +} + +impl<'a> OAuthService<'a> { + const SESS_STATE_FIELD: &'static str = "oauth_state"; + const SESS_PKCE_FIELD: &'static str = "oauth_verifier"; + + pub(super) fn new( + config: &'a AppConfig, + db: &'a DbPool, + http_client: &'a reqwest::Client, + ) -> Self { + Self { + config, + db, + http_client, + } + } + + pub async fn authorize_url( + &self, + provider: OAuthProviderEnum, + callback_path: &str, + session: &Session, + ) -> AppResult { + let provider = self.get_provider(provider)?; + let client = self.get_oauth_client(provider.as_ref(), callback_path)?; + + let state = oauth2::State::new_random(); + let pkce_verifier = oauth2::PkceCodeVerifierS256::new_random(); + let mut auth_url = client.authorize_url(&state); + auth_url + .query_pairs_mut() + .extend_pairs(pkce_verifier.authorize_url_params()); + + session.insert(Self::SESS_STATE_FIELD, state).await?; + session.insert(Self::SESS_PKCE_FIELD, pkce_verifier).await?; + + Ok(auth_url) + } + + pub async fn exchange_code( + &self, + provider: OAuthProviderEnum, + callback_path: &str, + session: &Session, + code: &str, + returned_state: &str, + ) -> AppResult { + // Get saved state and code verifier from session + let saved_state = session + .remove::(Self::SESS_STATE_FIELD) + .await? + .ok_or_else(|| AppError::unauthorized("no state in session"))?; + let pkce_verifier = session + .remove::(Self::SESS_PKCE_FIELD) + .await? + .ok_or_else(|| AppError::unauthorized("no PKCE verifier in session"))?; + + // Verify state + if saved_state.ct_ne(returned_state.as_bytes()).into() { + return Err(AppError::unauthorized("state parameter doesn't match")); + } + + // Exchange code for token + let provider = self.get_provider(provider)?; + let client = self.get_oauth_client(provider.as_ref(), callback_path)?; + let response = client + .exchange_code(code) + .param("code_verifier", String::from(pkce_verifier)) + .with_reqwest_client(&self.http_client) + .execute::() + .await + .map_err(|err| AppError::unauthorized(format!("token exchange failed: {err}")))?; + + Ok(response) + } + + pub async fn get_user( + &self, + provider: OAuthProviderEnum, + token: &oauth2::StandardToken, + active_session: Option, + ) -> AppResult { + let provider = self.get_provider(provider)?; + let mut user_info_request = self + .http_client + .get(provider.get_user_info_url()) + .bearer_auth(token.access_token().as_ref()); + for (name, value) in provider.create_request_headers() { + user_info_request = user_info_request.header(name, value); + } + let user_info_response = user_info_request.send().await.context("request failed")?; + if !user_info_response.status().is_success() { + let error = + anyhow::anyhow!("failed to get user: {:?}", user_info_response.text().await); + return Err(AppError::internal(error)); + } + let user_data = provider + .extract_user_data(user_info_response.json().await.context("request failed")?) + .context("unable to extract user data from response")?; + + let mut db = DbService::from_pool(self.db).await?; + let user = match provider.find_linked_user(&mut db, &user_data).await? { + Some(existing_user) => { + if active_session.is_some_and(|sess| sess.user_id != existing_user.id) { + return Err(AppError::unauthorized("cannot switch users via OAuth")); + } else { + existing_user + } + } + None => match active_session { + None => { + let new_user = provider.create_new_user(&user_data); + db.users().create(new_user).await? + } + Some(sess) => match db.users().find_by_id(&sess.user_id).await? { + Some(user) if provider.is_user_linked(&user) => { + return Err(AppError::bad_request("user already linked to provider")); + } + Some(user) => { + // Link logged-in user to new provider + let update_user = provider.create_update_user(&user_data); + db.users().update(&user.id, update_user).await?; + user + } + None => { + return Err(AppError::internal(anyhow::anyhow!("user not found"))); + } + }, + }, + }; + + Ok(user) + } + + fn get_provider(&self, provider: OAuthProviderEnum) -> AppResult> { + let provider: Option> = match provider { + OAuthProviderEnum::Github => match self.config.auth.github { + Some(ref c) => Some(Box::new(github::GitHubOAuthProvider::new(c))), + None => None, + }, + OAuthProviderEnum::Discord => match self.config.auth.discord { + Some(ref c) => Some(Box::new(discord::DiscordOAuthProvider::new(c))), + None => None, + }, + }; + + provider.ok_or_else(|| AppError::bad_request("unsupported OAuth provider")) + } + + fn get_oauth_client( + &self, + provider: &dyn OAuthProvider, + callback_path: &str, + ) -> anyhow::Result { + let mut client = oauth2::Client::new( + provider.get_client_id(), + provider.get_authorize_url().parse()?, + provider.get_token_url().parse()?, + ); + client.set_client_secret(provider.get_client_secret()); + client.set_redirect_url( + format!("{}{}", &self.config.server.base_url, callback_path).parse()?, + ); + for scope in provider.get_scopes() { + client.add_scope(scope); + } + + Ok(client) + } +} diff --git a/server-new/src/services/auth/oauth/service.rs b/server-new/src/services/auth/oauth/service.rs deleted file mode 100644 index e5614d6..0000000 --- a/server-new/src/services/auth/oauth/service.rs +++ /dev/null @@ -1,185 +0,0 @@ -use anyhow::Context; -use oauth2::{StandardToken, Token}; -use subtle::ConstantTimeEq; -use tower_sessions::Session; - -use crate::{ - config::AppConfig, - db::{DbPool, DbService, models::ChatRsUser}, - error::{AppError, AppResult}, - extractors::session::UserSession, - services::auth::oauth::{OAuthProvider, OAuthProviderEnum}, -}; - -/// OAuth functions -pub struct OAuthService<'a> { - config: &'a AppConfig, - db: &'a DbPool, - http_client: &'a reqwest::Client, -} - -impl<'a> OAuthService<'a> { - const SESS_STATE_FIELD: &'static str = "oauth_state"; - const SESS_PKCE_FIELD: &'static str = "oauth_verifier"; - - pub fn new(config: &'a AppConfig, db: &'a DbPool, http_client: &'a reqwest::Client) -> Self { - Self { - config, - db, - http_client, - } - } - - pub async fn authorize_url( - &self, - provider: OAuthProviderEnum, - callback_path: &str, - session: &Session, - ) -> AppResult { - let provider = self.get_provider(provider)?; - let client = self.get_oauth_client(provider.as_ref(), callback_path)?; - - let state = oauth2::State::new_random(); - let pkce_verifier = oauth2::PkceCodeVerifierS256::new_random(); - let mut auth_url = client.authorize_url(&state); - auth_url - .query_pairs_mut() - .extend_pairs(pkce_verifier.authorize_url_params()); - - session.insert(Self::SESS_STATE_FIELD, state).await?; - session.insert(Self::SESS_PKCE_FIELD, pkce_verifier).await?; - - Ok(auth_url) - } - - pub async fn exchange_code( - &self, - provider: OAuthProviderEnum, - callback_path: &str, - session: &Session, - code: &str, - returned_state: &str, - ) -> AppResult { - // Get saved state and code verifier from session - let saved_state = session - .remove::(Self::SESS_STATE_FIELD) - .await? - .ok_or_else(|| AppError::unauthorized("no state in session"))?; - let pkce_verifier = session - .remove::(Self::SESS_PKCE_FIELD) - .await? - .ok_or_else(|| AppError::unauthorized("no PKCE verifier in session"))?; - - // Verify state - if saved_state.ct_ne(returned_state.as_bytes()).into() { - return Err(AppError::unauthorized("state parameter doesn't match")); - } - - // Exchange code for token - let provider = self.get_provider(provider)?; - let client = self.get_oauth_client(provider.as_ref(), callback_path)?; - let response = client - .exchange_code(code) - .param("code_verifier", String::from(pkce_verifier)) - .with_reqwest_client(&self.http_client) - .execute::() - .await - .map_err(|err| AppError::unauthorized(format!("token exchange failed: {err}")))?; - - Ok(response) - } - - pub async fn get_user( - &self, - provider: OAuthProviderEnum, - token: &oauth2::StandardToken, - active_session: Option, - ) -> AppResult { - let provider = self.get_provider(provider)?; - let mut user_info_request = self - .http_client - .get(provider.get_user_info_url()) - .bearer_auth(token.access_token().as_ref()); - for (name, value) in provider.create_request_headers() { - user_info_request = user_info_request.header(name, value); - } - let user_info_response = user_info_request.send().await.context("request failed")?; - if !user_info_response.status().is_success() { - let error = - anyhow::anyhow!("failed to get user: {:?}", user_info_response.text().await); - return Err(AppError::internal(error)); - } - let user_data = provider - .extract_user_data(user_info_response.json().await.context("request failed")?) - .context("unable to extract user data from response")?; - - let mut db = DbService::from_pool(self.db).await?; - let user = match provider.find_linked_user(&mut db, &user_data).await? { - Some(existing_user) => { - if active_session.is_some_and(|sess| sess.user_id != existing_user.id) { - return Err(AppError::unauthorized("cannot switch users via OAuth")); - } else { - existing_user - } - } - None => match active_session { - None => { - let new_user = provider.create_new_user(&user_data); - db.users().create(new_user).await? - } - Some(sess) => match db.users().find_by_id(&sess.user_id).await? { - Some(user) if provider.is_user_linked(&user) => { - return Err(AppError::bad_request("user already linked to provider")); - } - Some(user) => { - // Link logged-in user to new provider - let update_user = provider.create_update_user(&user_data); - db.users().update(&user.id, update_user).await?; - user - } - None => { - return Err(AppError::internal(anyhow::anyhow!("user not found"))); - } - }, - }, - }; - - Ok(user) - } - - fn get_provider(&self, provider: OAuthProviderEnum) -> AppResult> { - let provider: Option> = match provider { - OAuthProviderEnum::Github => match self.config.auth.github { - Some(ref c) => Some(Box::new(super::github::GitHubOAuthProvider::new(c))), - None => None, - }, - OAuthProviderEnum::Discord => match self.config.auth.discord { - Some(ref c) => Some(Box::new(super::discord::DiscordOAuthProvider::new(c))), - None => None, - }, - }; - - provider.ok_or_else(|| AppError::bad_request("unsupported OAuth provider")) - } - - fn get_oauth_client( - &self, - provider: &dyn OAuthProvider, - callback_path: &str, - ) -> anyhow::Result { - let mut client = oauth2::Client::new( - provider.get_client_id(), - provider.get_authorize_url().parse()?, - provider.get_token_url().parse()?, - ); - client.set_client_secret(provider.get_client_secret()); - client.set_redirect_url( - format!("{}{}", &self.config.server.base_url, callback_path).parse()?, - ); - for scope in provider.get_scopes() { - client.add_scope(scope); - } - - Ok(client) - } -} diff --git a/server-new/src/services/auth/session.rs b/server-new/src/services/auth/session.rs new file mode 100644 index 0000000..9c6638c --- /dev/null +++ b/server-new/src/services/auth/session.rs @@ -0,0 +1,67 @@ +use std::collections::HashMap; + +use tower_sessions::{ + Expiry, Session, + cookie::time::Duration, + session_store::{Error as StoreError, Result as StoreResult}, +}; +use uuid::Uuid; + +use crate::{ + error::AppResult, + services::auth::types::{SessionMeta, UserSession}, +}; + +/// The field used to store the user ID in the session. +const USER_ID_FIELD: &str = "user_id"; +/// The field used to store the user session metadata. +const META_FIELD: &str = "meta"; + +/// Authentication-specific operations on a tower session. +pub struct AuthSessionService { + session_length: i64, +} + +impl AuthSessionService { + pub fn new(session_length: i64) -> Self { + Self { session_length } + } + + /// Initialize a new logged-in session for the given user. + pub async fn login( + &self, + session: &Session, + meta: &SessionMeta, + user_id: &Uuid, + ) -> AppResult<()> { + session.cycle_id().await?; + session.insert(USER_ID_FIELD, user_id).await?; + session.insert(META_FIELD, meta).await?; + session.set_expiry(Some(Expiry::OnInactivity(Duration::seconds( + self.session_length, + )))); + + Ok(()) + } + + /// Extract the current user session if this is an active user session. + pub async fn user_session(&self, session: &Session) -> AppResult> { + let user_id = session.get::(USER_ID_FIELD).await?; + Ok(user_id.map(UserSession::new)) + } + + /// Logout the user, deleting the current session. + pub async fn logout(&self, session: &Session) -> AppResult<()> { + session.flush().await?; + Ok(()) + } +} + +pub(super) fn user_id_from_record_data( + data: &HashMap, +) -> StoreResult> { + data.get(USER_ID_FIELD) + .map(|val| serde_json::from_value::(val.clone())) + .transpose() + .map_err(|_| StoreError::Encode("invalid user id field".to_owned())) +} diff --git a/server-new/src/services/auth/session_store.rs b/server-new/src/services/auth/session_store.rs index 31812e5..8608be4 100644 --- a/server-new/src/services/auth/session_store.rs +++ b/server-new/src/services/auth/session_store.rs @@ -23,15 +23,6 @@ impl SessionDbStore { Uuid::from_bytes(id.0.to_be_bytes()) } - fn extract_user_id(record: &Record) -> Result> { - record - .data - .get(super::USER_ID_FIELD) - .map(|val| serde_json::from_value::(val.clone())) - .transpose() - .map_err(|_| Error::Encode("invalid user id field".to_owned())) - } - fn convert_expiry(time: OffsetDateTime) -> Result { UtcDateTime::from_timestamp_secs(time.unix_timestamp()) .ok_or_else(|| Error::Backend(format!("Invalid expiry: {time}"))) @@ -55,7 +46,7 @@ impl SessionStore for SessionDbStore { /// Creates a new session in the store with the provided session record. async fn create(&self, record: &mut Record) -> Result<()> { let session_id = Self::get_session_uuid(&record.id); - let user_id = Self::extract_user_id(record)?; + let user_id = super::session::user_id_from_record_data(&record.data)?; let expires_at = Self::convert_expiry(record.expiry_date)?; let mut db = self.get_db().await?; diff --git a/server-new/src/services/auth/types.rs b/server-new/src/services/auth/types.rs new file mode 100644 index 0000000..c9206f6 --- /dev/null +++ b/server-new/src/services/auth/types.rs @@ -0,0 +1,26 @@ +use std::net::IpAddr; + +use serde::{Deserialize, Serialize}; +use uuid::Uuid; + +use crate::db::UtcDateTime; + +/// Active user session data. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct UserSession { + pub user_id: Uuid, +} + +impl UserSession { + pub fn new(user_id: Uuid) -> Self { + Self { user_id } + } +} + +/// Session metadata captured on login. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SessionMeta { + pub start_time: UtcDateTime, + pub ip: Option, + pub user_agent: Option, +} From 1d7a20082071ed59fd854da54fff0d4717819843 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 19:28:13 -0400 Subject: [PATCH 034/111] add google oauth --- server-new/src/config.rs | 3 +- server-new/src/db/repositories/user.rs | 18 ++-- server-new/src/services/auth/oauth/google.rs | 101 +++++++++++++++++++ server-new/src/services/auth/oauth/mod.rs | 9 ++ 4 files changed, 121 insertions(+), 10 deletions(-) create mode 100644 server-new/src/services/auth/oauth/google.rs diff --git a/server-new/src/config.rs b/server-new/src/config.rs index 8803f05..212cc8e 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -6,7 +6,7 @@ use figment::providers::{Env, Format, Toml}; use serde::Deserialize; use crate::{ - services::auth::oauth::{DiscordOAuthConfig, GitHubOAuthConfig}, + services::auth::oauth::{DiscordOAuthConfig, GitHubOAuthConfig, GoogleOAuthConfig}, state::AppState, }; @@ -42,6 +42,7 @@ pub struct AuthConfig { pub session_length: i64, pub github: Option, pub discord: Option, + pub google: Option, } #[derive(Debug, Clone, Deserialize)] diff --git a/server-new/src/db/repositories/user.rs b/server-new/src/db/repositories/user.rs index a1409bc..ad58f45 100644 --- a/server-new/src/db/repositories/user.rs +++ b/server-new/src/db/repositories/user.rs @@ -48,16 +48,16 @@ impl<'a> UserRepository<'a> { Ok(user) } - // pub async fn find_by_google_id(&mut self, id: &str) -> Result, Error> { - // let user = users::table - // .filter(users::google_id.eq(id)) - // .select(ChatRsUser::as_select()) - // .first(self.db) - // .await - // .optional()?; + pub async fn find_by_google_id(&mut self, id: &str) -> Result, Error> { + let user = users::table + .filter(users::google_id.eq(id)) + .select(ChatRsUser::as_select()) + .first(self.db) + .await + .optional()?; - // Ok(user) - // } + Ok(user) + } pub async fn find_by_discord_id(&mut self, id: &str) -> Result, Error> { let user = users::table diff --git a/server-new/src/services/auth/oauth/google.rs b/server-new/src/services/auth/oauth/google.rs new file mode 100644 index 0000000..00d101a --- /dev/null +++ b/server-new/src/services/auth/oauth/google.rs @@ -0,0 +1,101 @@ +use futures::future::BoxFuture; +use serde::Deserialize; + +use crate::{ + db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, + error::AppResult, + services::auth::oauth::{OAuthProvider, UserData}, +}; + +#[derive(Clone, Debug, Deserialize)] +pub struct GoogleOAuthConfig { + client_id: String, + client_secret: String, +} + +/// User info from Google API +#[derive(Debug, Deserialize)] +pub struct GoogleUserInfo { + sub: String, + name: String, + picture: Option, +} + +pub struct GoogleOAuthProvider { + config: GoogleOAuthConfig, +} + +impl GoogleOAuthProvider { + pub fn new(config: &GoogleOAuthConfig) -> Self { + Self { + config: config.clone(), + } + } +} + +impl OAuthProvider for GoogleOAuthProvider { + fn get_scopes(&self) -> Vec<&str> { + vec!["openid", "profile"] + } + + fn get_authorize_url(&self) -> &str { + "https://accounts.google.com/o/oauth2/v2/auth" + } + + fn get_token_url(&self) -> &str { + "https://oauth2.googleapis.com/token" + } + + fn get_user_info_url(&self) -> &str { + "https://www.googleapis.com/oauth2/v3/userinfo" + } + + fn get_client_id(&self) -> String { + self.config.client_id.clone() + } + + fn get_client_secret(&self) -> String { + self.config.client_secret.clone() + } + + fn extract_user_data(&self, user_info: serde_json::Value) -> anyhow::Result { + let user_info: GoogleUserInfo = serde_json::from_value(user_info)?; + + Ok(UserData { + id: user_info.sub, + name: user_info.name, + avatar_url: user_info.picture, + }) + } + + fn find_linked_user<'a>( + &self, + db: &'a mut crate::db::DbService, + user_data: &'a super::UserData, + ) -> BoxFuture<'a, AppResult>> { + Box::pin(async move { + let user = db.users().find_by_google_id(&user_data.id).await?; + Ok(user) + }) + } + + fn is_user_linked(&self, user: &ChatRsUser) -> bool { + user.google_id.is_some() + } + + fn create_update_user<'a>(&self, user_data: &'a super::UserData) -> UpdateChatRsUser<'a> { + UpdateChatRsUser { + google_id: Some(&user_data.id), + ..Default::default() + } + } + + fn create_new_user<'a>(&self, user_data: &'a super::UserData) -> NewChatRsUser<'a> { + NewChatRsUser { + google_id: Some(&user_data.id), + name: &user_data.name, + avatar_url: user_data.avatar_url.as_deref(), + ..Default::default() + } + } +} diff --git a/server-new/src/services/auth/oauth/mod.rs b/server-new/src/services/auth/oauth/mod.rs index 95427c5..d7c2bb7 100644 --- a/server-new/src/services/auth/oauth/mod.rs +++ b/server-new/src/services/auth/oauth/mod.rs @@ -14,11 +14,14 @@ use crate::{ error::{AppError, AppResult}, services::auth::UserSession, }; + mod discord; mod github; +mod google; pub use discord::DiscordOAuthConfig; pub use github::GitHubOAuthConfig; +pub use google::GoogleOAuthConfig; /// Supported OAuth providers #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] @@ -26,12 +29,14 @@ pub use github::GitHubOAuthConfig; pub enum OAuthProviderEnum { Github, Discord, + Google, } impl OAuthProviderEnum { pub fn as_str(&self) -> &str { match self { OAuthProviderEnum::Github => "github", OAuthProviderEnum::Discord => "discord", + OAuthProviderEnum::Google => "google", } } } @@ -216,6 +221,10 @@ impl<'a> OAuthService<'a> { Some(ref c) => Some(Box::new(discord::DiscordOAuthProvider::new(c))), None => None, }, + OAuthProviderEnum::Google => match self.config.auth.google { + Some(ref c) => Some(Box::new(google::GoogleOAuthProvider::new(c))), + None => None, + }, }; provider.ok_or_else(|| AppError::bad_request("unsupported OAuth provider")) From 3c6528074e9b9717e04dbfa9bcdacda73d2b5049 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 19:45:43 -0400 Subject: [PATCH 035/111] refactor oauth --- server-new/config.toml | 1 + server-new/src/services/auth/{oauth/mod.rs => oauth.rs} | 0 2 files changed, 1 insertion(+) rename server-new/src/services/auth/{oauth/mod.rs => oauth.rs} (100%) diff --git a/server-new/config.toml b/server-new/config.toml index 36f99fd..977cd6d 100644 --- a/server-new/config.toml +++ b/server-new/config.toml @@ -1,6 +1,7 @@ [server] host = "127.0.0.1" port = 8080 +base_url = "http://localhost:8080" log_level = "info" request_id_header = "x-request-id" diff --git a/server-new/src/services/auth/oauth/mod.rs b/server-new/src/services/auth/oauth.rs similarity index 100% rename from server-new/src/services/auth/oauth/mod.rs rename to server-new/src/services/auth/oauth.rs From 4e36b9883296aeaecbab5eae7d64ad8f46dedb2e Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 23 Jun 2026 23:09:44 -0400 Subject: [PATCH 036/111] refactor auth errors --- server-new/Cargo.lock | 1 + server-new/Cargo.toml | 1 + server-new/src/error.rs | 24 ------- server-new/src/extractors/session.rs | 7 +- server-new/src/services/auth/error.rs | 36 ++++++++++ server-new/src/services/auth/mod.rs | 7 +- server-new/src/services/auth/oauth.rs | 65 ++++++++++--------- server-new/src/services/auth/oauth/discord.rs | 11 ++-- server-new/src/services/auth/oauth/github.rs | 29 +++++---- server-new/src/services/auth/oauth/google.rs | 14 ++-- server-new/src/services/auth/session.rs | 12 ++-- 11 files changed, 119 insertions(+), 88 deletions(-) create mode 100644 server-new/src/services/auth/error.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index 42fd158..52dfba0 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -2088,6 +2088,7 @@ dependencies = [ "serde_json", "serde_with", "subtle", + "thiserror", "tokio", "tower", "tower-http 0.7.0", diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 9e62e0c..d5407a3 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -49,6 +49,7 @@ serde_with = { features = ["macros"] } subtle = "2.6.1" +thiserror = "2.0.18" tokio = { version = "1.52.3", default-features = false, diff --git a/server-new/src/error.rs b/server-new/src/error.rs index 9f1e5fd..dba225d 100644 --- a/server-new/src/error.rs +++ b/server-new/src/error.rs @@ -5,8 +5,6 @@ use axum::{ }; use serde::Serialize; -use crate::db::DbPoolError; - /// Global result type that can be used for API route handlers pub type AppResult = Result; @@ -18,28 +16,6 @@ pub struct AppError { source: Option, } -// Convenient error conversions -impl From for AppError { - fn from(error: anyhow::Error) -> Self { - Self::internal(error) - } -} -impl From for AppError { - fn from(error: diesel::result::Error) -> Self { - Self::internal(anyhow::Error::from(error).context("database error")) - } -} -impl From for AppError { - fn from(error: DbPoolError) -> Self { - Self::internal(anyhow::Error::from(error).context("pool error")) - } -} -impl From for AppError { - fn from(error: tower_sessions::session::Error) -> Self { - Self::internal(anyhow::Error::from(error).context("session error")) - } -} - impl AppError { pub fn new(status: StatusCode, message: impl Into) -> Self { Self { diff --git a/server-new/src/extractors/session.rs b/server-new/src/extractors/session.rs index ac9ef99..2293cd8 100644 --- a/server-new/src/extractors/session.rs +++ b/server-new/src/extractors/session.rs @@ -32,8 +32,13 @@ impl OptionalFromRequestParts for UserSession { let session = Session::from_request_parts(parts, state) .await .map_err(|(_, msg)| AppError::internal(anyhow::anyhow!(msg)))?; + let user_session = state + .auth_service() + .session() + .user_session(&session) + .await?; - state.auth_service().session().user_session(&session).await + Ok(user_session) } } diff --git a/server-new/src/services/auth/error.rs b/server-new/src/services/auth/error.rs new file mode 100644 index 0000000..a7a4a9b --- /dev/null +++ b/server-new/src/services/auth/error.rs @@ -0,0 +1,36 @@ +use crate::{db::DbPoolError, error::AppError}; + +pub type AuthResult = Result; + +/// Auth service errors +#[derive(Debug, thiserror::Error)] +pub enum AuthError { + #[error("{0}")] + Unauthorized(&'static str), + #[error("{0}")] + BadRequest(&'static str), + #[error(transparent)] + Provider(#[from] anyhow::Error), + #[error("user not found")] + UserNotFound, + #[error("database error: {0}")] + Database(#[from] diesel::result::Error), + #[error("database error: {0}")] + DatabasePool(#[from] DbPoolError), + #[error("session error: {0}")] + Session(#[from] tower_sessions::session::Error), + #[error("request error: {0}")] + Request(#[from] reqwest::Error), +} + +// Conversion to HTTP API errors +impl From for AppError { + fn from(error: AuthError) -> Self { + match error { + AuthError::Unauthorized(reason) => Self::unauthorized(reason), + AuthError::BadRequest(reason) => Self::bad_request(reason), + AuthError::Provider(error) => Self::internal(error.context("OAuth error")), + error => Self::internal(error.into()), + } + } +} diff --git a/server-new/src/services/auth/mod.rs b/server-new/src/services/auth/mod.rs index 8b0ffb3..f93dd63 100644 --- a/server-new/src/services/auth/mod.rs +++ b/server-new/src/services/auth/mod.rs @@ -1,15 +1,16 @@ use crate::{ config::AppConfig, db::{DbPool, DbService, models::ChatRsUser}, - error::{AppError, AppResult}, }; use uuid::Uuid; +mod error; pub mod oauth; pub mod session; pub mod session_store; mod types; +pub use error::{AuthError, AuthResult}; pub use types::*; pub struct AuthService<'a> { @@ -29,10 +30,10 @@ impl<'a> AuthService<'a> { /// Get the user from the database with the given ID, or return /// an internal error if not found - pub async fn get_user(&self, id: &Uuid) -> AppResult { + pub async fn get_user(&self, id: &Uuid) -> AuthResult { let mut db = DbService::from_pool(&self.db).await?; match db.users().find_by_id(id).await? { - None => Err(AppError::internal(anyhow::anyhow!("user not found"))), + None => Err(AuthError::UserNotFound), Some(user) => Ok(user), } } diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index d7c2bb7..a6f740a 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -1,6 +1,4 @@ -use anyhow::Context; use futures::future::BoxFuture; -use oauth2::{StandardToken, Token}; use serde::{Deserialize, Serialize}; use subtle::ConstantTimeEq; use tower_sessions::Session; @@ -11,8 +9,7 @@ use crate::{ DbPool, DbService, models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, }, - error::{AppError, AppResult}, - services::auth::UserSession, + services::auth::{AuthError, AuthResult, UserSession}, }; mod discord; @@ -52,12 +49,12 @@ pub trait OAuthProvider: Send + Sync { fn create_request_headers(&self) -> Vec<(&'static str, &'static str)> { vec![] } - fn extract_user_data(&self, user_info: serde_json::Value) -> anyhow::Result; + fn extract_user_data(&self, user_info: serde_json::Value) -> AuthResult; fn find_linked_user<'a>( &self, db: &'a mut DbService, user_data: &'a UserData, - ) -> BoxFuture<'a, AppResult>>; + ) -> BoxFuture<'a, AuthResult>>; fn is_user_linked(&self, user: &ChatRsUser) -> bool; fn create_update_user<'a>(&self, user_data: &'a UserData) -> UpdateChatRsUser<'a>; fn create_new_user<'a>(&self, user_data: &'a UserData) -> NewChatRsUser<'a>; @@ -71,6 +68,11 @@ pub struct UserData { pub avatar_url: Option, } +#[derive(Deserialize)] +pub struct TokenResponse { + access_token: String, +} + /// OAuth functions pub struct OAuthService<'a> { config: &'a AppConfig, @@ -99,7 +101,7 @@ impl<'a> OAuthService<'a> { provider: OAuthProviderEnum, callback_path: &str, session: &Session, - ) -> AppResult { + ) -> AuthResult { let provider = self.get_provider(provider)?; let client = self.get_oauth_client(provider.as_ref(), callback_path)?; @@ -122,21 +124,22 @@ impl<'a> OAuthService<'a> { callback_path: &str, session: &Session, code: &str, - returned_state: &str, - ) -> AppResult { + state: &str, + ) -> AuthResult { // Get saved state and code verifier from session let saved_state = session .remove::(Self::SESS_STATE_FIELD) .await? - .ok_or_else(|| AppError::unauthorized("no state in session"))?; + .ok_or(AuthError::Unauthorized("missing state in session"))? + .to_base64(); let pkce_verifier = session .remove::(Self::SESS_PKCE_FIELD) .await? - .ok_or_else(|| AppError::unauthorized("no PKCE verifier in session"))?; + .ok_or(AuthError::Unauthorized("missing PKCE in session"))?; // Verify state - if saved_state.ct_ne(returned_state.as_bytes()).into() { - return Err(AppError::unauthorized("state parameter doesn't match")); + if saved_state.as_bytes().ct_ne(state.as_bytes()).into() { + return Err(AuthError::Unauthorized("OAuth state mismatch")); } // Exchange code for token @@ -146,9 +149,11 @@ impl<'a> OAuthService<'a> { .exchange_code(code) .param("code_verifier", String::from(pkce_verifier)) .with_reqwest_client(&self.http_client) - .execute::() + .execute::() .await - .map_err(|err| AppError::unauthorized(format!("token exchange failed: {err}")))?; + .map_err(|err| { + AuthError::Provider(anyhow::Error::from(err).context("token exchange")) + })?; Ok(response) } @@ -156,32 +161,32 @@ impl<'a> OAuthService<'a> { pub async fn get_user( &self, provider: OAuthProviderEnum, - token: &oauth2::StandardToken, + token: &TokenResponse, active_session: Option, - ) -> AppResult { + ) -> AuthResult { let provider = self.get_provider(provider)?; + + // Get user info from provider let mut user_info_request = self .http_client .get(provider.get_user_info_url()) - .bearer_auth(token.access_token().as_ref()); + .bearer_auth(&token.access_token); for (name, value) in provider.create_request_headers() { user_info_request = user_info_request.header(name, value); } - let user_info_response = user_info_request.send().await.context("request failed")?; + let user_info_response = user_info_request.send().await?; if !user_info_response.status().is_success() { - let error = - anyhow::anyhow!("failed to get user: {:?}", user_info_response.text().await); - return Err(AppError::internal(error)); + let error = anyhow::anyhow!("couldn't get user: {:?}", user_info_response.text().await); + return Err(AuthError::Provider(error)); } - let user_data = provider - .extract_user_data(user_info_response.json().await.context("request failed")?) - .context("unable to extract user data from response")?; + let user_data = provider.extract_user_data(user_info_response.json().await?)?; + // Check for existing user, or create new user let mut db = DbService::from_pool(self.db).await?; let user = match provider.find_linked_user(&mut db, &user_data).await? { Some(existing_user) => { if active_session.is_some_and(|sess| sess.user_id != existing_user.id) { - return Err(AppError::unauthorized("cannot switch users via OAuth")); + return Err(AuthError::Unauthorized("cannot switch users via OAuth")); } else { existing_user } @@ -193,7 +198,7 @@ impl<'a> OAuthService<'a> { } Some(sess) => match db.users().find_by_id(&sess.user_id).await? { Some(user) if provider.is_user_linked(&user) => { - return Err(AppError::bad_request("user already linked to provider")); + return Err(AuthError::BadRequest("user already linked to provider")); } Some(user) => { // Link logged-in user to new provider @@ -202,7 +207,7 @@ impl<'a> OAuthService<'a> { user } None => { - return Err(AppError::internal(anyhow::anyhow!("user not found"))); + return Err(AuthError::UserNotFound); } }, }, @@ -211,7 +216,7 @@ impl<'a> OAuthService<'a> { Ok(user) } - fn get_provider(&self, provider: OAuthProviderEnum) -> AppResult> { + fn get_provider(&self, provider: OAuthProviderEnum) -> AuthResult> { let provider: Option> = match provider { OAuthProviderEnum::Github => match self.config.auth.github { Some(ref c) => Some(Box::new(github::GitHubOAuthProvider::new(c))), @@ -227,7 +232,7 @@ impl<'a> OAuthService<'a> { }, }; - provider.ok_or_else(|| AppError::bad_request("unsupported OAuth provider")) + provider.ok_or(AuthError::BadRequest("unsupported OAuth provider")) } fn get_oauth_client( diff --git a/server-new/src/services/auth/oauth/discord.rs b/server-new/src/services/auth/oauth/discord.rs index e7babbf..1f7fd75 100644 --- a/server-new/src/services/auth/oauth/discord.rs +++ b/server-new/src/services/auth/oauth/discord.rs @@ -3,8 +3,7 @@ use serde::Deserialize; use crate::{ db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, - error::AppResult, - services::auth::oauth::OAuthProvider, + services::auth::{AuthError, AuthResult, oauth::OAuthProvider}, }; #[derive(Clone, Debug, Deserialize)] @@ -50,8 +49,10 @@ impl OAuthProvider for DiscordOAuthProvider { self.config.client_secret.clone() } - fn extract_user_data(&self, user_info: serde_json::Value) -> anyhow::Result { - let user_info: DiscordUserInfo = serde_json::from_value(user_info)?; + fn extract_user_data(&self, user_info: serde_json::Value) -> AuthResult { + let user_info: DiscordUserInfo = serde_json::from_value(user_info).map_err(|err| { + AuthError::Provider(anyhow::Error::from(err).context("parse user info")) + })?; let avatar_url = user_info.avatar.as_ref().map(|avatar| { format!( "https://cdn.discordapp.com/avatars/{}/{}.png", @@ -70,7 +71,7 @@ impl OAuthProvider for DiscordOAuthProvider { &self, db: &'a mut crate::db::DbService, user_data: &'a super::UserData, - ) -> BoxFuture<'a, AppResult>> { + ) -> BoxFuture<'a, AuthResult>> { Box::pin(async move { let user = db.users().find_by_discord_id(&user_data.id).await?; Ok(user) diff --git a/server-new/src/services/auth/oauth/github.rs b/server-new/src/services/auth/oauth/github.rs index 65bc533..56dd804 100644 --- a/server-new/src/services/auth/oauth/github.rs +++ b/server-new/src/services/auth/oauth/github.rs @@ -3,8 +3,7 @@ use serde::Deserialize; use crate::{ db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, - error::AppResult, - services::auth::oauth::OAuthProvider, + services::auth::{AuthError, AuthResult, oauth::OAuthProvider}, }; #[derive(Clone, Debug, Deserialize)] @@ -13,6 +12,15 @@ pub struct GitHubOAuthConfig { client_secret: String, } +/// User info returned from GitHub API +#[derive(Debug, Deserialize)] +struct GitHubUserInfo { + id: u64, + login: String, + name: Option, + avatar_url: Option, +} + pub struct GitHubOAuthProvider { config: GitHubOAuthConfig, } @@ -57,8 +65,10 @@ impl OAuthProvider for GitHubOAuthProvider { ] } - fn extract_user_data(&self, user_info: serde_json::Value) -> anyhow::Result { - let info: GitHubUserInfo = serde_json::from_value(user_info)?; + fn extract_user_data(&self, user_info: serde_json::Value) -> AuthResult { + let info: GitHubUserInfo = serde_json::from_value(user_info).map_err(|err| { + AuthError::Provider(anyhow::Error::from(err).context("parse user info")) + })?; Ok(super::UserData { id: info.id.to_string(), @@ -71,7 +81,7 @@ impl OAuthProvider for GitHubOAuthProvider { &self, db: &'a mut crate::db::DbService, user_data: &'a super::UserData, - ) -> BoxFuture<'a, AppResult>> { + ) -> BoxFuture<'a, AuthResult>> { Box::pin(async move { let user = db.users().find_by_github_id(&user_data.id).await?; Ok(user) @@ -98,12 +108,3 @@ impl OAuthProvider for GitHubOAuthProvider { } } } - -/// User info returned from GitHub API -#[derive(Debug, Deserialize)] -struct GitHubUserInfo { - id: u64, - login: String, - name: Option, - avatar_url: Option, -} diff --git a/server-new/src/services/auth/oauth/google.rs b/server-new/src/services/auth/oauth/google.rs index 00d101a..19a2549 100644 --- a/server-new/src/services/auth/oauth/google.rs +++ b/server-new/src/services/auth/oauth/google.rs @@ -3,8 +3,10 @@ use serde::Deserialize; use crate::{ db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, - error::AppResult, - services::auth::oauth::{OAuthProvider, UserData}, + services::auth::{ + AuthError, AuthResult, + oauth::{OAuthProvider, UserData}, + }, }; #[derive(Clone, Debug, Deserialize)] @@ -58,8 +60,10 @@ impl OAuthProvider for GoogleOAuthProvider { self.config.client_secret.clone() } - fn extract_user_data(&self, user_info: serde_json::Value) -> anyhow::Result { - let user_info: GoogleUserInfo = serde_json::from_value(user_info)?; + fn extract_user_data(&self, user_info: serde_json::Value) -> AuthResult { + let user_info: GoogleUserInfo = serde_json::from_value(user_info).map_err(|err| { + AuthError::Provider(anyhow::Error::from(err).context("parse user info")) + })?; Ok(UserData { id: user_info.sub, @@ -72,7 +76,7 @@ impl OAuthProvider for GoogleOAuthProvider { &self, db: &'a mut crate::db::DbService, user_data: &'a super::UserData, - ) -> BoxFuture<'a, AppResult>> { + ) -> BoxFuture<'a, AuthResult>> { Box::pin(async move { let user = db.users().find_by_google_id(&user_data.id).await?; Ok(user) diff --git a/server-new/src/services/auth/session.rs b/server-new/src/services/auth/session.rs index 9c6638c..f7e61b8 100644 --- a/server-new/src/services/auth/session.rs +++ b/server-new/src/services/auth/session.rs @@ -7,9 +7,9 @@ use tower_sessions::{ }; use uuid::Uuid; -use crate::{ - error::AppResult, - services::auth::types::{SessionMeta, UserSession}, +use crate::services::auth::{ + AuthResult, + types::{SessionMeta, UserSession}, }; /// The field used to store the user ID in the session. @@ -33,7 +33,7 @@ impl AuthSessionService { session: &Session, meta: &SessionMeta, user_id: &Uuid, - ) -> AppResult<()> { + ) -> AuthResult<()> { session.cycle_id().await?; session.insert(USER_ID_FIELD, user_id).await?; session.insert(META_FIELD, meta).await?; @@ -45,13 +45,13 @@ impl AuthSessionService { } /// Extract the current user session if this is an active user session. - pub async fn user_session(&self, session: &Session) -> AppResult> { + pub async fn user_session(&self, session: &Session) -> AuthResult> { let user_id = session.get::(USER_ID_FIELD).await?; Ok(user_id.map(UserSession::new)) } /// Logout the user, deleting the current session. - pub async fn logout(&self, session: &Session) -> AppResult<()> { + pub async fn logout(&self, session: &Session) -> AuthResult<()> { session.flush().await?; Ok(()) } From 700f97fd62903c19e1ec88cb03e83efd9719ba50 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 24 Jun 2026 20:40:13 -0400 Subject: [PATCH 037/111] Update config.rs --- server-new/src/config.rs | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/server-new/src/config.rs b/server-new/src/config.rs index 212cc8e..a0aa96d 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -60,7 +60,11 @@ pub struct RedisConfig { pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Config").on_init(async |mut state| { let config = extract_config()?; - tracing::info!(log_level = config.server.log_level, "Config loaded!"); + tracing::info!( + log_level = config.server.log_level, + base_url = config.server.base_url, + "Config loaded!" + ); state.insert(config); Ok(state) }) From 4ac1654ff2e91891cbf20d90c31d009618d4a22d Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sat, 27 Jun 2026 03:15:29 -0400 Subject: [PATCH 038/111] extract oauth library --- server-new/Cargo.lock | 316 +++++------------- server-new/Cargo.toml | 10 +- server-new/crates/simple-oauth/Cargo.toml | 17 + server-new/crates/simple-oauth/src/common.rs | 3 + .../crates/simple-oauth/src/common/discord.rs | 68 ++++ .../crates/simple-oauth/src/common/github.rs | 78 +++++ .../crates/simple-oauth/src/common/google.rs | 61 ++++ server-new/crates/simple-oauth/src/lib.rs | 130 +++++++ .../crates/simple-oauth/src/provider.rs | 18 + server-new/crates/simple-oauth/src/types.rs | 17 + server-new/src/services/auth/error.rs | 7 +- server-new/src/services/auth/oauth.rs | 134 +++----- server-new/src/services/auth/oauth/discord.rs | 63 +--- server-new/src/services/auth/oauth/github.rs | 65 +--- server-new/src/services/auth/oauth/google.rs | 59 +--- 15 files changed, 558 insertions(+), 488 deletions(-) create mode 100644 server-new/crates/simple-oauth/Cargo.toml create mode 100644 server-new/crates/simple-oauth/src/common.rs create mode 100644 server-new/crates/simple-oauth/src/common/discord.rs create mode 100644 server-new/crates/simple-oauth/src/common/github.rs create mode 100644 server-new/crates/simple-oauth/src/common/google.rs create mode 100644 server-new/crates/simple-oauth/src/lib.rs create mode 100644 server-new/crates/simple-oauth/src/provider.rs create mode 100644 server-new/crates/simple-oauth/src/types.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index 52dfba0..aa6a989 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -8,7 +8,7 @@ version = "0.5.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0" dependencies = [ - "crypto-common 0.1.7", + "crypto-common", "generic-array", ] @@ -20,7 +20,7 @@ checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0" dependencies = [ "cfg-if", "cipher", - "cpufeatures 0.2.17", + "cpufeatures", ] [[package]] @@ -70,24 +70,6 @@ dependencies = [ "rustversion", ] -[[package]] -name = "async-oauth2" -version = "0.6.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b82d800aa8c98755d56f8492f8861f6f958e750e1ce940f55846e94d4defef28" -dependencies = [ - "base64", - "bytes", - "http", - "rand 0.10.1", - "reqwest", - "serde", - "serde-aux", - "serde_json", - "sha2 0.11.0", - "url", -] - [[package]] name = "async-trait" version = "0.1.89" @@ -250,15 +232,6 @@ dependencies = [ "generic-array", ] -[[package]] -name = "block-buffer" -version = "0.12.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d2f6c7dbe95a6ed67ad9f18e57daf93a2f034c524b99fd2b76d18fdfeb6660aa" -dependencies = [ - "hybrid-array", -] - [[package]] name = "bumpalo" version = "3.20.3" @@ -317,17 +290,6 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" -[[package]] -name = "chacha20" -version = "0.10.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6f8d983286843e49675a4b7a2d174efe136dc93a18d69130dd18198a6c167601" -dependencies = [ - "cfg-if", - "cpufeatures 0.3.0", - "rand_core 0.10.1", -] - [[package]] name = "chrono" version = "0.4.45" @@ -335,8 +297,10 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1aa79e62e7697b8e29b513a68abacf485adcd1fe8284a4316c5ae868e6633327" dependencies = [ "iana-time-zone", + "js-sys", "num-traits", "serde", + "wasm-bindgen", "windows-link", ] @@ -346,7 +310,7 @@ version = "0.4.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" dependencies = [ - "crypto-common 0.1.7", + "crypto-common", "inout", ] @@ -369,12 +333,6 @@ dependencies = [ "memchr", ] -[[package]] -name = "const-oid" -version = "0.10.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a6ef517f0926dd24a1582492c791b6a4818a4d94e789a334894aa15b0d12f55c" - [[package]] name = "cookie" version = "0.18.1" @@ -387,7 +345,7 @@ dependencies = [ "hmac", "percent-encoding", "rand 0.8.6", - "sha2 0.10.9", + "sha2", "subtle", "time", "version_check", @@ -399,16 +357,6 @@ version = "0.3.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "396de984970346b0d9e93d1415082923c679e5ae5c3ee3dcbd104f5610af126b" -[[package]] -name = "core-foundation" -version = "0.9.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "91e195e091a93c46f7102ec7818a2aa394e1e1771c3ab4825963fa03e45afb8f" -dependencies = [ - "core-foundation-sys", - "libc", -] - [[package]] name = "core-foundation" version = "0.10.1" @@ -434,15 +382,6 @@ dependencies = [ "libc", ] -[[package]] -name = "cpufeatures" -version = "0.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" -dependencies = [ - "libc", -] - [[package]] name = "crc16" version = "0.4.0" @@ -475,15 +414,6 @@ dependencies = [ "typenum", ] -[[package]] -name = "crypto-common" -version = "0.2.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ce6e4c961d6cd6c9a86db418387425e8bdeaf05b3c8bc1411e6dca4c252f1453" -dependencies = [ - "hybrid-array", -] - [[package]] name = "ctr" version = "0.9.2" @@ -668,22 +598,11 @@ version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ - "block-buffer 0.10.4", - "crypto-common 0.1.7", + "block-buffer", + "crypto-common", "subtle", ] -[[package]] -name = "digest" -version = "0.11.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f1dd6dbb5841937940781866fa1281a1ff7bd3bf827091440879f9994983d5c2" -dependencies = [ - "block-buffer 0.12.1", - "const-oid", - "crypto-common 0.2.2", -] - [[package]] name = "displaydoc" version = "0.2.6" @@ -733,15 +652,6 @@ version = "1.16.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" -[[package]] -name = "encoding_rs" -version = "0.8.35" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "75030f3c4f45dafd7586dd6780965a8c7e8e285a5ecb86713e63a79c5b2766f3" -dependencies = [ - "cfg-if", -] - [[package]] name = "equivalent" version = "1.0.2" @@ -985,7 +895,6 @@ dependencies = [ "cfg-if", "libc", "r-efi 6.0.0", - "rand_core 0.10.1", ] [[package]] @@ -998,25 +907,6 @@ dependencies = [ "polyval", ] -[[package]] -name = "h2" -version = "0.4.15" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6cb093c84e8bd9b188d4c4a8cb6579fc016968d14c99882163cd3ff402a4f155" -dependencies = [ - "atomic-waker", - "bytes", - "fnv", - "futures-core", - "futures-sink", - "http", - "indexmap", - "slab", - "tokio", - "tokio-util", - "tracing", -] - [[package]] name = "hashbrown" version = "0.17.1" @@ -1062,7 +952,7 @@ version = "0.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" dependencies = [ - "digest 0.10.7", + "digest", ] [[package]] @@ -1110,15 +1000,6 @@ version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" -[[package]] -name = "hybrid-array" -version = "0.4.12" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9155a582abd142abc056962c29e3ce5ff2ad5469f4246b537ed42c5deba857da" -dependencies = [ - "typenum", -] - [[package]] name = "hyper" version = "1.10.1" @@ -1129,7 +1010,6 @@ dependencies = [ "bytes", "futures-channel", "futures-core", - "h2", "http", "http-body", "httparse", @@ -1174,11 +1054,9 @@ dependencies = [ "percent-encoding", "pin-project-lite", "socket2 0.6.4", - "system-configuration", "tokio", "tower-service", "tracing", - "windows-registry", ] [[package]] @@ -1363,7 +1241,7 @@ dependencies = [ "jni-sys", "log", "simd_cesu8", - "thiserror", + "thiserror 2.0.18", "walkdir", "windows-link", ] @@ -1491,7 +1369,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d89e7ee0cfbedfc4da3340218492196241d89eefb6dab27de5df917a6d2e78cf" dependencies = [ "cfg-if", - "digest 0.10.7", + "digest", ] [[package]] @@ -1588,6 +1466,35 @@ dependencies = [ "libc", ] +[[package]] +name = "oauth2" +version = "5.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "51e219e79014df21a225b1860a479e2dcd7cbd9130f4defd4bd0e191ea31d67d" +dependencies = [ + "base64", + "chrono", + "getrandom 0.2.17", + "http", + "rand 0.8.6", + "serde", + "serde_json", + "serde_path_to_error", + "sha2", + "thiserror 1.0.69", + "url", +] + +[[package]] +name = "oauth2-reqwest" +version = "0.1.0-alpha.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "234fb5c965bbce983ee5de636a7a51d6a3223da8067ea02f9ab2d2d78ac08be2" +dependencies = [ + "oauth2", + "reqwest", +] + [[package]] name = "objc2-core-foundation" version = "0.3.2" @@ -1624,15 +1531,6 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe" -[[package]] -name = "ordered-float" -version = "2.10.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "68f19d67e5a2795c94e73e0bb1cc1a7edeb2e28efd39e2e1c9b7a40c1108b11c" -dependencies = [ - "num-traits", -] - [[package]] name = "parking_lot" version = "0.12.5" @@ -1717,7 +1615,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25" dependencies = [ "cfg-if", - "cpufeatures 0.2.17", + "cpufeatures", "opaque-debug", "universal-hash", ] @@ -1736,7 +1634,7 @@ dependencies = [ "md-5", "memchr", "rand 0.9.4", - "sha2 0.10.9", + "sha2", "stringprep", ] @@ -1811,7 +1709,7 @@ dependencies = [ "rustc-hash", "rustls", "socket2 0.6.4", - "thiserror", + "thiserror 2.0.18", "tokio", "tracing", "web-time", @@ -1833,7 +1731,7 @@ dependencies = [ "rustls", "rustls-pki-types", "slab", - "thiserror", + "thiserror 2.0.18", "tinyvec", "tracing", "web-time", @@ -1895,17 +1793,6 @@ dependencies = [ "rand_core 0.9.5", ] -[[package]] -name = "rand" -version = "0.10.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d2e8e8bcc7961af1fdac401278c6a831614941f6164ee3bf4ce61b7edb162207" -dependencies = [ - "chacha20", - "getrandom 0.4.3", - "rand_core 0.10.1", -] - [[package]] name = "rand_chacha" version = "0.3.1" @@ -1944,12 +1831,6 @@ dependencies = [ "getrandom 0.3.4", ] -[[package]] -name = "rand_core" -version = "0.10.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" - [[package]] name = "redis-protocol" version = "6.0.0" @@ -1998,9 +1879,7 @@ checksum = "219c5811de6525e5416c7d5d53bb656d3afdbc6c5af816e0802bcfa42dbdc1c3" dependencies = [ "base64", "bytes", - "encoding_rs", "futures-core", - "h2", "http", "http-body", "http-body-util", @@ -2009,7 +1888,6 @@ dependencies = [ "hyper-util", "js-sys", "log", - "mime", "percent-encoding", "pin-project-lite", "quinn", @@ -2068,7 +1946,6 @@ name = "rs-chat-api" version = "0.1.0" dependencies = [ "anyhow", - "async-oauth2", "async-trait", "axum", "axum-helmet", @@ -2087,8 +1964,9 @@ dependencies = [ "serde", "serde_json", "serde_with", + "simple-oauth", "subtle", - "thiserror", + "thiserror 2.0.18", "tokio", "tower", "tower-http 0.7.0", @@ -2157,7 +2035,7 @@ version = "0.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "26d1e2536ce4f35f4846aa13bff16bd0ff40157cdb14cc056c7b14ba41233ba0" dependencies = [ - "core-foundation 0.10.1", + "core-foundation", "core-foundation-sys", "jni", "log", @@ -2233,7 +2111,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d" dependencies = [ "bitflags", - "core-foundation 0.10.1", + "core-foundation", "core-foundation-sys", "libc", "security-framework-sys", @@ -2265,28 +2143,6 @@ dependencies = [ "serde_derive", ] -[[package]] -name = "serde-aux" -version = "4.7.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "207f67b28fe90fb596503a9bf0bf1ea5e831e21307658e177c5dfcdfc3ab8a0a" -dependencies = [ - "chrono", - "serde", - "serde-value", - "serde_json", -] - -[[package]] -name = "serde-value" -version = "0.7.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f3a1a3341211875ef120e117ea7fd5228530ae7e7036a779fdc9117be6b3282c" -dependencies = [ - "ordered-float", - "serde", -] - [[package]] name = "serde_core" version = "1.0.228" @@ -2390,19 +2246,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" dependencies = [ "cfg-if", - "cpufeatures 0.2.17", - "digest 0.10.7", -] - -[[package]] -name = "sha2" -version = "0.11.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "446ba717509524cb3f22f17ecc096f10f4822d76ab5c0b9822c5f9c284e825f4" -dependencies = [ - "cfg-if", - "cpufeatures 0.3.0", - "digest 0.11.3", + "cpufeatures", + "digest", ] [[package]] @@ -2446,6 +2291,18 @@ version = "0.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" +[[package]] +name = "simple-oauth" +version = "0.1.0" +dependencies = [ + "oauth2", + "oauth2-reqwest", + "reqwest", + "serde", + "serde_json", + "thiserror 2.0.18", +] + [[package]] name = "siphasher" version = "1.0.2" @@ -2551,33 +2408,32 @@ dependencies = [ ] [[package]] -name = "system-configuration" -version = "0.7.0" +name = "thiserror" +version = "1.0.69" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a13f3d0daba03132c0aa9767f98351b3488edc2c100cda2d2ec2b04f3d8d3c8b" +checksum = "b6aaf5339b578ea85b50e080feb250a3e8ae8cfcdff9a461c9ec2904bc923f52" dependencies = [ - "bitflags", - "core-foundation 0.9.4", - "system-configuration-sys", + "thiserror-impl 1.0.69", ] [[package]] -name = "system-configuration-sys" -version = "0.6.0" +name = "thiserror" +version = "2.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8e1d1b10ced5ca923a1fcb8d03e96b8d3268065d724548c0211415ff6ac6bac4" +checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" dependencies = [ - "core-foundation-sys", - "libc", + "thiserror-impl 2.0.18", ] [[package]] -name = "thiserror" -version = "2.0.18" +name = "thiserror-impl" +version = "1.0.69" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" +checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1" dependencies = [ - "thiserror-impl", + "proc-macro2", + "quote", + "syn", ] [[package]] @@ -2929,7 +2785,7 @@ dependencies = [ "rand 0.9.4", "serde", "serde_json", - "thiserror", + "thiserror 2.0.18", "time", "tokio", "tracing", @@ -2955,7 +2811,7 @@ dependencies = [ "async-trait", "fred", "rmp-serde", - "thiserror", + "thiserror 2.0.18", "time", "tower-sessions-core", ] @@ -2980,7 +2836,7 @@ checksum = "050686193eb999b4bb3bc2acfa891a13da00f79734704c4b8b4ef1a10b368a3c" dependencies = [ "crossbeam-channel", "symlink", - "thiserror", + "thiserror 2.0.18", "time", "tracing-subscriber", ] @@ -3111,7 +2967,7 @@ version = "0.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea" dependencies = [ - "crypto-common 0.1.7", + "crypto-common", "subtle", ] @@ -3131,6 +2987,7 @@ dependencies = [ "idna", "percent-encoding", "serde", + "serde_derive", ] [[package]] @@ -3368,17 +3225,6 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" -[[package]] -name = "windows-registry" -version = "0.6.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "02752bf7fbdcce7f2a27a742f798510f3e5ad88dbe84871e5168e2120c3d5720" -dependencies = [ - "windows-link", - "windows-result", - "windows-strings", -] - [[package]] name = "windows-result" version = "0.4.1" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index d5407a3..dc876a7 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -7,7 +7,6 @@ publish = false [dependencies] anyhow = "1.0.102" -async-oauth2 = "0.6.0" async-trait = "0.1.89" axum = { version = "0.8.9", features = ["json", "query"] } axum-helmet = "1.0.2" @@ -29,7 +28,7 @@ diesel-async = { version = "0.9.2", features = ["deadpool", "migrations", "postgres"] } -diesel-jsonb-derive = { path = "./crates/diesel-jsonb-derive" } +diesel-jsonb-derive = { path = "crates/diesel-jsonb-derive" } diesel_migrations = { version = "2.3.2", features = ["postgres"] } dotenvy = "0.15.7" figment = { version = "0.10.19", features = ["env", "toml"] } @@ -40,7 +39,11 @@ fred = { } futures = "0.3.32" hex = "0.4.3" -reqwest = { version = "0.13.4", features = ["json"] } +reqwest = { + version = "0.13.4", + default-features = false, + features = ["default-tls", "json"] +} serde = { version = "1.0.228", features = ["derive"] } serde_json = "1.0.150" serde_with = { @@ -48,6 +51,7 @@ serde_with = { default-features = false, features = ["macros"] } +simple-oauth = { path = "crates/simple-oauth/", features = ["default-tls"] } subtle = "2.6.1" thiserror = "2.0.18" tokio = { diff --git a/server-new/crates/simple-oauth/Cargo.toml b/server-new/crates/simple-oauth/Cargo.toml new file mode 100644 index 0000000..fc1e9a4 --- /dev/null +++ b/server-new/crates/simple-oauth/Cargo.toml @@ -0,0 +1,17 @@ +[package] +name = "simple-oauth" +version = "0.1.0" +edition = "2024" + +[dependencies] +oauth2 = { version = "5", default-features = false } +oauth2-reqwest = "0.1.0-alpha.3" +reqwest = { version = "0.13", default-features = false, features = ["json"] } +serde = "1.0" +serde_json = "1.0" +thiserror = "2.0" + +[features] +default = ["default-tls"] +default-tls = ["reqwest/default-tls"] +native-tls = ["reqwest/native-tls"] diff --git a/server-new/crates/simple-oauth/src/common.rs b/server-new/crates/simple-oauth/src/common.rs new file mode 100644 index 0000000..51106db --- /dev/null +++ b/server-new/crates/simple-oauth/src/common.rs @@ -0,0 +1,3 @@ +pub mod discord; +pub mod github; +pub mod google; diff --git a/server-new/crates/simple-oauth/src/common/discord.rs b/server-new/crates/simple-oauth/src/common/discord.rs new file mode 100644 index 0000000..b96e436 --- /dev/null +++ b/server-new/crates/simple-oauth/src/common/discord.rs @@ -0,0 +1,68 @@ +use serde::Deserialize; + +use crate::{SimpleOAuthProvider, types::UserInfo}; + +pub struct Discord { + client_id: u64, + client_secret: String, +} + +impl Discord { + pub fn new(client_id: u64, client_secret: impl Into) -> Self { + Self { + client_id, + client_secret: client_secret.into(), + } + } +} + +/// User info returned from Discord API +#[derive(Debug, Deserialize)] +pub struct DiscordUserInfo { + id: String, + username: String, + global_name: Option, + avatar: Option, +} + +impl SimpleOAuthProvider for Discord { + fn get_authorize_url(&self) -> String { + "https://discord.com/oauth2/authorize".to_owned() + } + + fn get_token_url(&self) -> String { + "https://discord.com/api/oauth2/token".to_owned() + } + + fn get_scopes(&self) -> Vec { + vec!["identify".to_owned()] + } + + fn get_user_info_url(&self) -> String { + "https://discord.com/api/v9/users/@me".to_owned() + } + + fn get_client_id(&self) -> String { + self.client_id.to_string() + } + + fn get_client_secret(&self) -> String { + self.client_secret.clone() + } + + fn extract_user_info(&self, val: serde_json::Value) -> Result { + let user_info: DiscordUserInfo = serde_json::from_value(val)?; + let avatar_url = user_info.avatar.as_ref().map(|avatar| { + format!( + "https://cdn.discordapp.com/avatars/{}/{}.png", + user_info.id, avatar + ) + }); + + Ok(UserInfo { + id: user_info.id, + name: user_info.global_name.unwrap_or_else(|| user_info.username), + avatar_url, + }) + } +} diff --git a/server-new/crates/simple-oauth/src/common/github.rs b/server-new/crates/simple-oauth/src/common/github.rs new file mode 100644 index 0000000..a8cd198 --- /dev/null +++ b/server-new/crates/simple-oauth/src/common/github.rs @@ -0,0 +1,78 @@ +use serde::Deserialize; + +use crate::{SimpleOAuthProvider, types::UserInfo}; + +pub struct GitHub { + client_id: String, + client_secret: String, + user_agent: String, +} + +impl GitHub { + pub fn new( + client_id: impl Into, + client_secret: impl Into, + user_agent: impl Into, + ) -> Self { + Self { + client_id: client_id.into(), + client_secret: client_secret.into(), + user_agent: user_agent.into(), + } + } +} + +/// User info returned from GitHub API +#[derive(Debug, Deserialize)] +struct GitHubUserInfo { + id: u64, + login: String, + name: Option, + avatar_url: Option, +} + +impl SimpleOAuthProvider for GitHub { + fn get_authorize_url(&self) -> String { + String::from("https://github.com/login/oauth/authorize") + } + + fn get_token_url(&self) -> String { + String::from("https://github.com/login/oauth/access_token") + } + + fn get_scopes(&self) -> Vec { + vec!["user:read".into()] + } + + fn get_user_info_url(&self) -> String { + String::from("https://api.github.com/user") + } + + fn get_client_id(&self) -> String { + self.client_id.to_owned() + } + + fn get_client_secret(&self) -> String { + self.client_secret.to_owned() + } + + fn create_request_headers(&self) -> Vec<(String, String)> { + vec![ + ("Accept".into(), "application/vnd.github+json".into()), + ("User-Agent".into(), self.user_agent.clone()), + ] + } + + fn extract_user_info( + &self, + user_info: serde_json::Value, + ) -> Result { + let info: GitHubUserInfo = serde_json::from_value(user_info)?; + + Ok(UserInfo { + id: info.id.to_string(), + name: info.name.unwrap_or(info.login), + avatar_url: info.avatar_url, + }) + } +} diff --git a/server-new/crates/simple-oauth/src/common/google.rs b/server-new/crates/simple-oauth/src/common/google.rs new file mode 100644 index 0000000..7bbd93e --- /dev/null +++ b/server-new/crates/simple-oauth/src/common/google.rs @@ -0,0 +1,61 @@ +use serde::Deserialize; + +use crate::{SimpleOAuthProvider, types::UserInfo}; + +pub struct Google { + client_id: String, + client_secret: String, +} + +impl Google { + pub fn new(client_id: impl Into, client_secret: impl Into) -> Self { + Self { + client_id: client_id.into(), + client_secret: client_secret.into(), + } + } +} + +/// User info from Google API +#[derive(Debug, Deserialize)] +pub struct GoogleUserInfo { + sub: String, + name: String, + picture: Option, +} + +impl SimpleOAuthProvider for Google { + fn get_scopes(&self) -> Vec { + vec!["openid".into(), "profile".into()] + } + + fn get_authorize_url(&self) -> String { + "https://accounts.google.com/o/oauth2/v2/auth".into() + } + + fn get_token_url(&self) -> String { + "https://oauth2.googleapis.com/token".into() + } + + fn get_user_info_url(&self) -> String { + "https://www.googleapis.com/oauth2/v3/userinfo".into() + } + + fn get_client_id(&self) -> String { + self.client_id.clone() + } + + fn get_client_secret(&self) -> String { + self.client_secret.clone() + } + + fn extract_user_info(&self, val: serde_json::Value) -> Result { + let user_info: GoogleUserInfo = serde_json::from_value(val)?; + + Ok(UserInfo { + id: user_info.sub, + name: user_info.name, + avatar_url: user_info.picture, + }) + } +} diff --git a/server-new/crates/simple-oauth/src/lib.rs b/server-new/crates/simple-oauth/src/lib.rs new file mode 100644 index 0000000..85af483 --- /dev/null +++ b/server-new/crates/simple-oauth/src/lib.rs @@ -0,0 +1,130 @@ +use oauth2::{ + CsrfToken, HttpClientError, RequestTokenError, TokenResponse, basic::BasicErrorResponse, +}; + +pub mod common; +mod provider; +pub mod types; + +pub use provider::SimpleOAuthProvider; + +use crate::types::{AuthorizeUrl, StandardTokenResponse, UserInfo}; + +#[derive(Debug, thiserror::Error)] +pub enum SimpleOAuthError { + #[error(transparent)] + Request(#[from] reqwest::Error), + #[error("invalid url: {0}")] + ParseUrl(#[from] oauth2::url::ParseError), + #[error("token exchange error: {0}")] + TokenExchange(#[from] RequestTokenError, BasicErrorResponse>), + #[error("deserialization error: {0}")] + Deserialization(#[from] serde_json::Error), +} + +pub struct SimpleOAuthClient { + http_client: reqwest::Client, + oauth_client: oauth2_reqwest::ReqwestClient, +} + +impl SimpleOAuthClient { + pub fn new() -> Result { + let http_client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build()?; + + Ok(Self { + oauth_client: oauth2_reqwest::ReqwestClient::from(http_client.clone()), + http_client, + }) + } + + pub fn with_http_client(http_client: reqwest::Client) -> Self { + Self { + oauth_client: oauth2_reqwest::ReqwestClient::from(http_client.clone()), + http_client, + } + } + + pub fn authorize_url( + &self, + provider: &P, + redirect_url: &str, + scopes: Option>, + ) -> Result { + let oauth_client = + oauth2::basic::BasicClient::new(oauth2::ClientId::new(provider.get_client_id())) + .set_client_secret(oauth2::ClientSecret::new(provider.get_client_secret())) + .set_auth_uri(oauth2::AuthUrl::new(provider.get_authorize_url())?) + .set_redirect_uri(oauth2::RedirectUrl::new(redirect_url.into())?); + let (pkce_challenge, pkce_verifier) = oauth2::PkceCodeChallenge::new_random_sha256(); + let (url, state) = oauth_client + .authorize_url(CsrfToken::new_random) + .add_scopes( + scopes + .map(|scopes| scopes.into_iter().map(|s| s.to_owned()).collect()) + .unwrap_or_else(|| provider.get_scopes()) + .into_iter() + .map(oauth2::Scope::new), + ) + .set_pkce_challenge(pkce_challenge) + .url(); + + Ok(AuthorizeUrl { + url, + state: state.into_secret(), + pkce_verifier: pkce_verifier.into_secret(), + }) + } + + pub async fn exchange_code( + &self, + provider: &P, + redirect_url: &str, + code: &str, + pkce_verifier: Option<&str>, + ) -> Result { + let oauth_client = + oauth2::basic::BasicClient::new(oauth2::ClientId::new(provider.get_client_id())) + .set_client_secret(oauth2::ClientSecret::new(provider.get_client_secret())) + .set_redirect_uri(oauth2::RedirectUrl::new(redirect_url.into())?) + .set_token_uri(oauth2::TokenUrl::new(provider.get_token_url())?); + let mut token_request = + oauth_client.exchange_code(oauth2::AuthorizationCode::new(code.into())); + if let Some(verifier) = pkce_verifier { + token_request = + token_request.set_pkce_verifier(oauth2::PkceCodeVerifier::new(verifier.into())); + } + let token = token_request.request_async(&self.oauth_client).await?; + + Ok(StandardTokenResponse { + access_token: token.access_token().secret().to_owned(), + refresh_token: token.refresh_token().map(|s| s.secret().to_owned()), + expires_in: token.expires_in(), + }) + } + + pub async fn get_user_info( + &self, + provider: &P, + access_token: &str, + ) -> Result { + let mut user_info_request = self + .http_client + .get(provider.get_user_info_url()) + .bearer_auth(access_token); + for (name, val) in provider.create_request_headers() { + user_info_request = user_info_request.header(name, val); + } + + let user_info_val = user_info_request + .send() + .await? + .error_for_status()? + .json() + .await?; + let user_info = provider.extract_user_info(user_info_val)?; + + Ok(user_info) + } +} diff --git a/server-new/crates/simple-oauth/src/provider.rs b/server-new/crates/simple-oauth/src/provider.rs new file mode 100644 index 0000000..a33b900 --- /dev/null +++ b/server-new/crates/simple-oauth/src/provider.rs @@ -0,0 +1,18 @@ +use crate::types::UserInfo; + +/// Trait for all OAuth providers +pub trait SimpleOAuthProvider: Send + Sync { + fn get_scopes(&self) -> Vec; + fn get_authorize_url(&self) -> String; + fn get_token_url(&self) -> String; + fn get_user_info_url(&self) -> String; + fn get_client_id(&self) -> String; + fn get_client_secret(&self) -> String; + fn create_request_headers(&self) -> Vec<(String, String)> { + vec![] + } + fn extract_user_info( + &self, + user_info: serde_json::Value, + ) -> Result; +} diff --git a/server-new/crates/simple-oauth/src/types.rs b/server-new/crates/simple-oauth/src/types.rs new file mode 100644 index 0000000..89d4d8f --- /dev/null +++ b/server-new/crates/simple-oauth/src/types.rs @@ -0,0 +1,17 @@ +pub struct AuthorizeUrl { + pub url: oauth2::url::Url, + pub state: String, + pub pkce_verifier: String, +} + +pub struct StandardTokenResponse { + pub access_token: String, + pub refresh_token: Option, + pub expires_in: Option, +} + +pub struct UserInfo { + pub id: String, + pub name: String, + pub avatar_url: Option, +} diff --git a/server-new/src/services/auth/error.rs b/server-new/src/services/auth/error.rs index a7a4a9b..b2e7854 100644 --- a/server-new/src/services/auth/error.rs +++ b/server-new/src/services/auth/error.rs @@ -9,8 +9,8 @@ pub enum AuthError { Unauthorized(&'static str), #[error("{0}")] BadRequest(&'static str), - #[error(transparent)] - Provider(#[from] anyhow::Error), + #[error("OAuth error: {0}")] + OAuth(#[from] simple_oauth::SimpleOAuthError), #[error("user not found")] UserNotFound, #[error("database error: {0}")] @@ -19,8 +19,6 @@ pub enum AuthError { DatabasePool(#[from] DbPoolError), #[error("session error: {0}")] Session(#[from] tower_sessions::session::Error), - #[error("request error: {0}")] - Request(#[from] reqwest::Error), } // Conversion to HTTP API errors @@ -29,7 +27,6 @@ impl From for AppError { match error { AuthError::Unauthorized(reason) => Self::unauthorized(reason), AuthError::BadRequest(reason) => Self::bad_request(reason), - AuthError::Provider(error) => Self::internal(error.context("OAuth error")), error => Self::internal(error.into()), } } diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index a6f740a..927f152 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -1,5 +1,9 @@ use futures::future::BoxFuture; use serde::{Deserialize, Serialize}; +use simple_oauth::{ + SimpleOAuthProvider, + types::{AuthorizeUrl, StandardTokenResponse, UserInfo}, +}; use subtle::ConstantTimeEq; use tower_sessions::Session; @@ -40,37 +44,15 @@ impl OAuthProviderEnum { /// Trait for all OAuth providers pub trait OAuthProvider: Send + Sync { - fn get_scopes(&self) -> Vec<&str>; - fn get_authorize_url(&self) -> &str; - fn get_token_url(&self) -> &str; - fn get_user_info_url(&self) -> &str; - fn get_client_id(&self) -> String; - fn get_client_secret(&self) -> String; - fn create_request_headers(&self) -> Vec<(&'static str, &'static str)> { - vec![] - } - fn extract_user_data(&self, user_info: serde_json::Value) -> AuthResult; + fn get_inner_provider(&self) -> Box; fn find_linked_user<'a>( &self, db: &'a mut DbService, - user_data: &'a UserData, + user_info: &'a UserInfo, ) -> BoxFuture<'a, AuthResult>>; fn is_user_linked(&self, user: &ChatRsUser) -> bool; - fn create_update_user<'a>(&self, user_data: &'a UserData) -> UpdateChatRsUser<'a>; - fn create_new_user<'a>(&self, user_data: &'a UserData) -> NewChatRsUser<'a>; -} - -/// Common OAuth user data structure -#[derive(Debug)] -pub struct UserData { - pub id: String, - pub name: String, - pub avatar_url: Option, -} - -#[derive(Deserialize)] -pub struct TokenResponse { - access_token: String, + fn create_update_user<'a>(&self, user_info: &'a UserInfo) -> UpdateChatRsUser<'a>; + fn create_new_user<'a>(&self, user_info: &'a UserInfo) -> NewChatRsUser<'a>; } /// OAuth functions @@ -101,21 +83,23 @@ impl<'a> OAuthService<'a> { provider: OAuthProviderEnum, callback_path: &str, session: &Session, - ) -> AuthResult { + ) -> AuthResult { let provider = self.get_provider(provider)?; - let client = self.get_oauth_client(provider.as_ref(), callback_path)?; - - let state = oauth2::State::new_random(); - let pkce_verifier = oauth2::PkceCodeVerifierS256::new_random(); - let mut auth_url = client.authorize_url(&state); - auth_url - .query_pairs_mut() - .extend_pairs(pkce_verifier.authorize_url_params()); - + let client = simple_oauth::SimpleOAuthClient::with_http_client(self.http_client.clone()); + + let AuthorizeUrl { + url, + state, + pkce_verifier, + } = client.authorize_url( + provider.get_inner_provider().as_ref(), + &self.get_redirect_url(callback_path), + None, + )?; session.insert(Self::SESS_STATE_FIELD, state).await?; session.insert(Self::SESS_PKCE_FIELD, pkce_verifier).await?; - Ok(auth_url) + Ok(url) } pub async fn exchange_code( @@ -125,15 +109,14 @@ impl<'a> OAuthService<'a> { session: &Session, code: &str, state: &str, - ) -> AuthResult { + ) -> AuthResult { // Get saved state and code verifier from session let saved_state = session - .remove::(Self::SESS_STATE_FIELD) + .remove::(Self::SESS_STATE_FIELD) .await? - .ok_or(AuthError::Unauthorized("missing state in session"))? - .to_base64(); + .ok_or(AuthError::Unauthorized("missing state in session"))?; let pkce_verifier = session - .remove::(Self::SESS_PKCE_FIELD) + .remove::(Self::SESS_PKCE_FIELD) .await? .ok_or(AuthError::Unauthorized("missing PKCE in session"))?; @@ -144,16 +127,15 @@ impl<'a> OAuthService<'a> { // Exchange code for token let provider = self.get_provider(provider)?; - let client = self.get_oauth_client(provider.as_ref(), callback_path)?; + let client = simple_oauth::SimpleOAuthClient::with_http_client(self.http_client.clone()); let response = client - .exchange_code(code) - .param("code_verifier", String::from(pkce_verifier)) - .with_reqwest_client(&self.http_client) - .execute::() - .await - .map_err(|err| { - AuthError::Provider(anyhow::Error::from(err).context("token exchange")) - })?; + .exchange_code( + provider.get_inner_provider().as_ref(), + &self.get_redirect_url(callback_path), + code, + Some(&pkce_verifier), + ) + .await?; Ok(response) } @@ -161,29 +143,20 @@ impl<'a> OAuthService<'a> { pub async fn get_user( &self, provider: OAuthProviderEnum, - token: &TokenResponse, + token: &StandardTokenResponse, active_session: Option, ) -> AuthResult { let provider = self.get_provider(provider)?; + let client = simple_oauth::SimpleOAuthClient::with_http_client(self.http_client.clone()); // Get user info from provider - let mut user_info_request = self - .http_client - .get(provider.get_user_info_url()) - .bearer_auth(&token.access_token); - for (name, value) in provider.create_request_headers() { - user_info_request = user_info_request.header(name, value); - } - let user_info_response = user_info_request.send().await?; - if !user_info_response.status().is_success() { - let error = anyhow::anyhow!("couldn't get user: {:?}", user_info_response.text().await); - return Err(AuthError::Provider(error)); - } - let user_data = provider.extract_user_data(user_info_response.json().await?)?; + let user_info = client + .get_user_info(provider.get_inner_provider().as_ref(), &token.access_token) + .await?; // Check for existing user, or create new user let mut db = DbService::from_pool(self.db).await?; - let user = match provider.find_linked_user(&mut db, &user_data).await? { + let user = match provider.find_linked_user(&mut db, &user_info).await? { Some(existing_user) => { if active_session.is_some_and(|sess| sess.user_id != existing_user.id) { return Err(AuthError::Unauthorized("cannot switch users via OAuth")); @@ -193,7 +166,7 @@ impl<'a> OAuthService<'a> { } None => match active_session { None => { - let new_user = provider.create_new_user(&user_data); + let new_user = provider.create_new_user(&user_info); db.users().create(new_user).await? } Some(sess) => match db.users().find_by_id(&sess.user_id).await? { @@ -202,7 +175,7 @@ impl<'a> OAuthService<'a> { } Some(user) => { // Link logged-in user to new provider - let update_user = provider.create_update_user(&user_data); + let update_user = provider.create_update_user(&user_info); db.users().update(&user.id, update_user).await?; user } @@ -216,6 +189,10 @@ impl<'a> OAuthService<'a> { Ok(user) } + fn get_redirect_url(&self, callback_path: &str) -> String { + format!("{}{}", &self.config.server.base_url, callback_path) + } + fn get_provider(&self, provider: OAuthProviderEnum) -> AuthResult> { let provider: Option> = match provider { OAuthProviderEnum::Github => match self.config.auth.github { @@ -234,25 +211,4 @@ impl<'a> OAuthService<'a> { provider.ok_or(AuthError::BadRequest("unsupported OAuth provider")) } - - fn get_oauth_client( - &self, - provider: &dyn OAuthProvider, - callback_path: &str, - ) -> anyhow::Result { - let mut client = oauth2::Client::new( - provider.get_client_id(), - provider.get_authorize_url().parse()?, - provider.get_token_url().parse()?, - ); - client.set_client_secret(provider.get_client_secret()); - client.set_redirect_url( - format!("{}{}", &self.config.server.base_url, callback_path).parse()?, - ); - for scope in provider.get_scopes() { - client.add_scope(scope); - } - - Ok(client) - } } diff --git a/server-new/src/services/auth/oauth/discord.rs b/server-new/src/services/auth/oauth/discord.rs index 1f7fd75..d593464 100644 --- a/server-new/src/services/auth/oauth/discord.rs +++ b/server-new/src/services/auth/oauth/discord.rs @@ -1,9 +1,10 @@ use futures::future::BoxFuture; use serde::Deserialize; +use simple_oauth::{common::discord::Discord, types::UserInfo}; use crate::{ db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, - services::auth::{AuthError, AuthResult, oauth::OAuthProvider}, + services::auth::{AuthResult, oauth::OAuthProvider}, }; #[derive(Clone, Debug, Deserialize)] @@ -25,52 +26,17 @@ impl DiscordOAuthProvider { } impl OAuthProvider for DiscordOAuthProvider { - fn get_authorize_url(&self) -> &str { - "https://discord.com/oauth2/authorize" - } - - fn get_token_url(&self) -> &str { - "https://discord.com/api/oauth2/token" - } - - fn get_scopes(&self) -> Vec<&str> { - vec!["identify"] - } - - fn get_user_info_url(&self) -> &str { - "https://discord.com/api/v9/users/@me" - } - - fn get_client_id(&self) -> String { - self.config.client_id.to_string() - } - - fn get_client_secret(&self) -> String { - self.config.client_secret.clone() - } - - fn extract_user_data(&self, user_info: serde_json::Value) -> AuthResult { - let user_info: DiscordUserInfo = serde_json::from_value(user_info).map_err(|err| { - AuthError::Provider(anyhow::Error::from(err).context("parse user info")) - })?; - let avatar_url = user_info.avatar.as_ref().map(|avatar| { - format!( - "https://cdn.discordapp.com/avatars/{}/{}.png", - user_info.id, avatar - ) - }); - - Ok(super::UserData { - id: user_info.id, - name: user_info.global_name.unwrap_or_else(|| user_info.username), - avatar_url, - }) + fn get_inner_provider(&self) -> Box { + Box::new(Discord::new( + self.config.client_id, + &self.config.client_secret, + )) } fn find_linked_user<'a>( &self, db: &'a mut crate::db::DbService, - user_data: &'a super::UserData, + user_data: &'a UserInfo, ) -> BoxFuture<'a, AuthResult>> { Box::pin(async move { let user = db.users().find_by_discord_id(&user_data.id).await?; @@ -82,14 +48,14 @@ impl OAuthProvider for DiscordOAuthProvider { user.discord_id.is_some() } - fn create_update_user<'a>(&self, user_data: &'a super::UserData) -> UpdateChatRsUser<'a> { + fn create_update_user<'a>(&self, user_data: &'a UserInfo) -> UpdateChatRsUser<'a> { UpdateChatRsUser { discord_id: Some(&user_data.id), ..Default::default() } } - fn create_new_user<'a>(&self, user_data: &'a super::UserData) -> NewChatRsUser<'a> { + fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { discord_id: Some(&user_data.id), name: &user_data.name, @@ -98,12 +64,3 @@ impl OAuthProvider for DiscordOAuthProvider { } } } - -/// User info returned from Discord API -#[derive(Debug, Deserialize)] -pub struct DiscordUserInfo { - id: String, - username: String, - global_name: Option, - avatar: Option, -} diff --git a/server-new/src/services/auth/oauth/github.rs b/server-new/src/services/auth/oauth/github.rs index 56dd804..15d976b 100644 --- a/server-new/src/services/auth/oauth/github.rs +++ b/server-new/src/services/auth/oauth/github.rs @@ -1,9 +1,10 @@ use futures::future::BoxFuture; use serde::Deserialize; +use simple_oauth::{SimpleOAuthProvider, common::github::GitHub, types::UserInfo}; use crate::{ db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, - services::auth::{AuthError, AuthResult, oauth::OAuthProvider}, + services::auth::{AuthResult, oauth::OAuthProvider}, }; #[derive(Clone, Debug, Deserialize)] @@ -12,15 +13,6 @@ pub struct GitHubOAuthConfig { client_secret: String, } -/// User info returned from GitHub API -#[derive(Debug, Deserialize)] -struct GitHubUserInfo { - id: u64, - login: String, - name: Option, - avatar_url: Option, -} - pub struct GitHubOAuthProvider { config: GitHubOAuthConfig, } @@ -34,53 +26,18 @@ impl GitHubOAuthProvider { } impl OAuthProvider for GitHubOAuthProvider { - fn get_authorize_url(&self) -> &str { - "https://github.com/login/oauth/authorize" - } - - fn get_token_url(&self) -> &str { - "https://github.com/login/oauth/access_token" - } - - fn get_scopes(&self) -> Vec<&str> { - vec!["user:read"] - } - - fn get_user_info_url(&self) -> &str { - "https://api.github.com/user" - } - - fn get_client_id(&self) -> String { - self.config.client_id.to_owned() - } - - fn get_client_secret(&self) -> String { - self.config.client_secret.to_owned() - } - - fn create_request_headers(&self) -> Vec<(&'static str, &'static str)> { - vec![ - ("Accept", "application/vnd.github+json"), - ("User-Agent", "fa-sharp/rs-chat"), - ] - } - - fn extract_user_data(&self, user_info: serde_json::Value) -> AuthResult { - let info: GitHubUserInfo = serde_json::from_value(user_info).map_err(|err| { - AuthError::Provider(anyhow::Error::from(err).context("parse user info")) - })?; - - Ok(super::UserData { - id: info.id.to_string(), - name: info.name.unwrap_or(info.login), - avatar_url: info.avatar_url, - }) + fn get_inner_provider(&self) -> Box { + Box::new(GitHub::new( + &self.config.client_id, + &self.config.client_secret, + "fa-sharp/rs-chat", + )) } fn find_linked_user<'a>( &self, db: &'a mut crate::db::DbService, - user_data: &'a super::UserData, + user_data: &'a UserInfo, ) -> BoxFuture<'a, AuthResult>> { Box::pin(async move { let user = db.users().find_by_github_id(&user_data.id).await?; @@ -92,14 +49,14 @@ impl OAuthProvider for GitHubOAuthProvider { user.github_id.is_some() } - fn create_update_user<'a>(&self, user_data: &'a super::UserData) -> UpdateChatRsUser<'a> { + fn create_update_user<'a>(&self, user_data: &'a UserInfo) -> UpdateChatRsUser<'a> { UpdateChatRsUser { github_id: Some(&user_data.id), ..Default::default() } } - fn create_new_user<'a>(&self, user_data: &'a super::UserData) -> NewChatRsUser<'a> { + fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { github_id: Some(&user_data.id), name: &user_data.name, diff --git a/server-new/src/services/auth/oauth/google.rs b/server-new/src/services/auth/oauth/google.rs index 19a2549..f77b043 100644 --- a/server-new/src/services/auth/oauth/google.rs +++ b/server-new/src/services/auth/oauth/google.rs @@ -1,12 +1,10 @@ use futures::future::BoxFuture; use serde::Deserialize; +use simple_oauth::{common::google::Google, types::UserInfo}; use crate::{ db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, - services::auth::{ - AuthError, AuthResult, - oauth::{OAuthProvider, UserData}, - }, + services::auth::{AuthResult, oauth::OAuthProvider}, }; #[derive(Clone, Debug, Deserialize)] @@ -15,14 +13,6 @@ pub struct GoogleOAuthConfig { client_secret: String, } -/// User info from Google API -#[derive(Debug, Deserialize)] -pub struct GoogleUserInfo { - sub: String, - name: String, - picture: Option, -} - pub struct GoogleOAuthProvider { config: GoogleOAuthConfig, } @@ -36,46 +26,17 @@ impl GoogleOAuthProvider { } impl OAuthProvider for GoogleOAuthProvider { - fn get_scopes(&self) -> Vec<&str> { - vec!["openid", "profile"] - } - - fn get_authorize_url(&self) -> &str { - "https://accounts.google.com/o/oauth2/v2/auth" - } - - fn get_token_url(&self) -> &str { - "https://oauth2.googleapis.com/token" - } - - fn get_user_info_url(&self) -> &str { - "https://www.googleapis.com/oauth2/v3/userinfo" - } - - fn get_client_id(&self) -> String { - self.config.client_id.clone() - } - - fn get_client_secret(&self) -> String { - self.config.client_secret.clone() - } - - fn extract_user_data(&self, user_info: serde_json::Value) -> AuthResult { - let user_info: GoogleUserInfo = serde_json::from_value(user_info).map_err(|err| { - AuthError::Provider(anyhow::Error::from(err).context("parse user info")) - })?; - - Ok(UserData { - id: user_info.sub, - name: user_info.name, - avatar_url: user_info.picture, - }) + fn get_inner_provider(&self) -> Box { + Box::new(Google::new( + &self.config.client_id, + &self.config.client_secret, + )) } fn find_linked_user<'a>( &self, db: &'a mut crate::db::DbService, - user_data: &'a super::UserData, + user_data: &'a UserInfo, ) -> BoxFuture<'a, AuthResult>> { Box::pin(async move { let user = db.users().find_by_google_id(&user_data.id).await?; @@ -87,14 +48,14 @@ impl OAuthProvider for GoogleOAuthProvider { user.google_id.is_some() } - fn create_update_user<'a>(&self, user_data: &'a super::UserData) -> UpdateChatRsUser<'a> { + fn create_update_user<'a>(&self, user_data: &'a UserInfo) -> UpdateChatRsUser<'a> { UpdateChatRsUser { google_id: Some(&user_data.id), ..Default::default() } } - fn create_new_user<'a>(&self, user_data: &'a super::UserData) -> NewChatRsUser<'a> { + fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { google_id: Some(&user_data.id), name: &user_data.name, From 469cbc885f0e7dcbb2f1226d4def07594b3da5eb Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sat, 27 Jun 2026 03:57:49 -0400 Subject: [PATCH 039/111] session cleanup task --- server-new/src/db/repositories/session.rs | 9 ++- server-new/src/plugins/session.rs | 75 ++++++++++++++++------- server-new/src/services/auth/session.rs | 21 +++++-- 3 files changed, 75 insertions(+), 30 deletions(-) diff --git a/server-new/src/db/repositories/session.rs b/server-new/src/db/repositories/session.rs index 6e14206..a4cfa98 100644 --- a/server-new/src/db/repositories/session.rs +++ b/server-new/src/db/repositories/session.rs @@ -71,7 +71,14 @@ impl<'a> SessionRepository<'a> { /// Delete a session by ID. Won't return an error if it does not exist. pub async fn delete_by_id(&mut self, session_id: &Uuid) -> QueryResult { diesel::delete(auth_sessions::table.find(session_id)) - .returning(auth_sessions::id) + .execute(self.db) + .await + } + + /// Delete all expired sessions + pub async fn delete_expired(&mut self) -> QueryResult { + diesel::delete(auth_sessions::table) + .filter(auth_sessions::expires_at.le(diesel::dsl::now)) .execute(self.db) .await } diff --git a/server-new/src/plugins/session.rs b/server-new/src/plugins/session.rs index 176af99..6f111b5 100644 --- a/server-new/src/plugins/session.rs +++ b/server-new/src/plugins/session.rs @@ -1,3 +1,5 @@ +use std::time::Duration; + use anyhow::{Context, bail}; use axum_plugin::AdHocPlugin; use tower_sessions::{ @@ -6,32 +8,59 @@ use tower_sessions::{ }; use tower_sessions_redis_store::RedisStore; -use crate::{services::auth::session_store::SessionDbStore, state::AppState}; +use crate::{ + db::DbPool, + services::auth::{session::AuthSessionService, session_store::SessionDbStore}, + state::AppState, +}; const REDIS_PREFIX: &str = "rs-chat:sess:"; +const CLEANUP_INTERVAL: Duration = Duration::from_mins(15); /// Add session handling to the server. Sessions are stored in Postgres and cached in Redis. pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("Session").on_setup(|router, state: &AppState| { - let cookie_key = - hex::decode(&state.config.auth.cookie_key).context("cookie_key must be hex value")?; - if cookie_key.len() < 32 { - bail!("cookie_key must be at least 32 bytes"); - } - - let redis_store = RedisStore::with_prefix(state.redis.clone(), REDIS_PREFIX.to_owned()); - let db_store = SessionDbStore::new(state.db_pool.clone()); - let session_store = CachingSessionStore::new(redis_store, db_store); - - let session_layer = SessionManagerLayer::new(session_store) - .with_name(state.config.auth.cookie_name.clone()) - .with_expiry(Expiry::OnSessionEnd) - .with_private(Key::derive_from(&cookie_key)) - .with_path("/") - .with_secure(true) - .with_http_only(true) - .with_same_site(SameSite::Lax); - - Ok(router.layer(session_layer)) - }) + AdHocPlugin::named("Session") + .on_init(async |state| { + let db_pool = state.get::().context("no db pool")?.clone(); + + // Session cleanup task + tokio::task::spawn(async move { + let mut interval = tokio::time::interval(CLEANUP_INTERVAL); + interval.tick().await; + + loop { + interval.tick().await; + tracing::debug!("Cleaning up auth sessions"); + if let Err(err) = AuthSessionService::session_cleanup(&db_pool).await { + tracing::warn!("Error cleaning up auth sessions: {err}"); + } + } + }); + + Ok(state) + }) + .on_setup(|router, state: &AppState| { + let cookie_key = hex::decode(&state.config.auth.cookie_key) + .context("cookie_key must be hex value")?; + if cookie_key.len() < 32 { + bail!("cookie_key must be at least 32 bytes"); + } + + // Session persistence + let redis_store = RedisStore::with_prefix(state.redis.clone(), REDIS_PREFIX.to_owned()); + let db_store = SessionDbStore::new(state.db_pool.clone()); + let session_store = CachingSessionStore::new(redis_store, db_store); + + // Add session / cookie management to router + let session_layer = SessionManagerLayer::new(session_store) + .with_name(state.config.auth.cookie_name.clone()) + .with_expiry(Expiry::OnSessionEnd) + .with_private(Key::derive_from(&cookie_key)) + .with_path("/") + .with_secure(true) + .with_http_only(true) + .with_same_site(SameSite::Lax); + + Ok(router.layer(session_layer)) + }) } diff --git a/server-new/src/services/auth/session.rs b/server-new/src/services/auth/session.rs index f7e61b8..3ff3102 100644 --- a/server-new/src/services/auth/session.rs +++ b/server-new/src/services/auth/session.rs @@ -7,9 +7,12 @@ use tower_sessions::{ }; use uuid::Uuid; -use crate::services::auth::{ - AuthResult, - types::{SessionMeta, UserSession}, +use crate::{ + db::{DbPool, DbService}, + services::auth::{ + AuthResult, + types::{SessionMeta, UserSession}, + }, }; /// The field used to store the user ID in the session. @@ -34,7 +37,7 @@ impl AuthSessionService { meta: &SessionMeta, user_id: &Uuid, ) -> AuthResult<()> { - session.cycle_id().await?; + session.cycle_id().await?; // ensures that the user id is saved to the database session.insert(USER_ID_FIELD, user_id).await?; session.insert(META_FIELD, meta).await?; session.set_expiry(Some(Expiry::OnInactivity(Duration::seconds( @@ -52,8 +55,14 @@ impl AuthSessionService { /// Logout the user, deleting the current session. pub async fn logout(&self, session: &Session) -> AuthResult<()> { - session.flush().await?; - Ok(()) + Ok(session.flush().await?) + } + + // Cleanup expired sessions + #[tracing::instrument(skip(db_pool), level = "debug")] + pub async fn session_cleanup(db_pool: &DbPool) -> AuthResult { + let mut db = DbService::from_pool(&db_pool).await?; + Ok(db.sessions().delete_expired().await?) } } From 42e33f3ca8366d37efb4ed451f647e3ea0185bc5 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sat, 27 Jun 2026 19:55:42 -0400 Subject: [PATCH 040/111] add oidc, moar oauth tweaks --- server-new/crates/simple-oauth/Cargo.toml | 1 + server-new/crates/simple-oauth/src/common.rs | 2 + .../crates/simple-oauth/src/common/discord.rs | 47 ++++------ .../crates/simple-oauth/src/common/github.rs | 52 ++++-------- .../crates/simple-oauth/src/common/google.rs | 50 ++++------- .../crates/simple-oauth/src/common/oidc.rs | 85 +++++++++++++++++++ server-new/crates/simple-oauth/src/lib.rs | 30 ++++--- .../crates/simple-oauth/src/provider.rs | 19 ++--- server-new/crates/simple-oauth/src/types.rs | 74 +++++++++++++++- server-new/src/services/auth/oauth.rs | 14 +-- server-new/src/services/auth/oauth/discord.rs | 14 +-- server-new/src/services/auth/oauth/github.rs | 17 ++-- server-new/src/services/auth/oauth/google.rs | 13 +-- 13 files changed, 267 insertions(+), 151 deletions(-) create mode 100644 server-new/crates/simple-oauth/src/common/oidc.rs diff --git a/server-new/crates/simple-oauth/Cargo.toml b/server-new/crates/simple-oauth/Cargo.toml index fc1e9a4..99c721b 100644 --- a/server-new/crates/simple-oauth/Cargo.toml +++ b/server-new/crates/simple-oauth/Cargo.toml @@ -2,6 +2,7 @@ name = "simple-oauth" version = "0.1.0" edition = "2024" +description = "Simple OAuth2 authorization crate" [dependencies] oauth2 = { version = "5", default-features = false } diff --git a/server-new/crates/simple-oauth/src/common.rs b/server-new/crates/simple-oauth/src/common.rs index 51106db..844d508 100644 --- a/server-new/crates/simple-oauth/src/common.rs +++ b/server-new/crates/simple-oauth/src/common.rs @@ -1,3 +1,5 @@ pub mod discord; pub mod github; pub mod google; + +pub mod oidc; diff --git a/server-new/crates/simple-oauth/src/common/discord.rs b/server-new/crates/simple-oauth/src/common/discord.rs index b96e436..119a1cb 100644 --- a/server-new/crates/simple-oauth/src/common/discord.rs +++ b/server-new/crates/simple-oauth/src/common/discord.rs @@ -2,52 +2,35 @@ use serde::Deserialize; use crate::{SimpleOAuthProvider, types::UserInfo}; -pub struct Discord { - client_id: u64, - client_secret: String, -} - -impl Discord { - pub fn new(client_id: u64, client_secret: impl Into) -> Self { - Self { - client_id, - client_secret: client_secret.into(), - } - } -} +#[derive(Debug)] +pub struct Discord; /// User info returned from Discord API #[derive(Debug, Deserialize)] -pub struct DiscordUserInfo { +struct DiscordUserInfo { id: String, username: String, global_name: Option, + email: Option, + verified: Option, avatar: Option, } impl SimpleOAuthProvider for Discord { - fn get_authorize_url(&self) -> String { - "https://discord.com/oauth2/authorize".to_owned() - } - - fn get_token_url(&self) -> String { - "https://discord.com/api/oauth2/token".to_owned() - } - - fn get_scopes(&self) -> Vec { - vec!["identify".to_owned()] + fn authorize_url(&self) -> &str { + "https://discord.com/oauth2/authorize" } - fn get_user_info_url(&self) -> String { - "https://discord.com/api/v9/users/@me".to_owned() + fn token_url(&self) -> &str { + "https://discord.com/api/oauth2/token" } - fn get_client_id(&self) -> String { - self.client_id.to_string() + fn default_scopes(&self) -> Vec<&str> { + vec!["identify"] } - fn get_client_secret(&self) -> String { - self.client_secret.clone() + fn user_info_url(&self) -> &str { + "https://discord.com/api/v9/users/@me" } fn extract_user_info(&self, val: serde_json::Value) -> Result { @@ -61,7 +44,9 @@ impl SimpleOAuthProvider for Discord { Ok(UserInfo { id: user_info.id, - name: user_info.global_name.unwrap_or_else(|| user_info.username), + email: user_info.email, + email_verified: user_info.verified, + name: user_info.global_name.or(Some(user_info.username)), avatar_url, }) } diff --git a/server-new/crates/simple-oauth/src/common/github.rs b/server-new/crates/simple-oauth/src/common/github.rs index a8cd198..a5b1119 100644 --- a/server-new/crates/simple-oauth/src/common/github.rs +++ b/server-new/crates/simple-oauth/src/common/github.rs @@ -2,25 +2,8 @@ use serde::Deserialize; use crate::{SimpleOAuthProvider, types::UserInfo}; -pub struct GitHub { - client_id: String, - client_secret: String, - user_agent: String, -} - -impl GitHub { - pub fn new( - client_id: impl Into, - client_secret: impl Into, - user_agent: impl Into, - ) -> Self { - Self { - client_id: client_id.into(), - client_secret: client_secret.into(), - user_agent: user_agent.into(), - } - } -} +#[derive(Debug)] +pub struct GitHub; /// User info returned from GitHub API #[derive(Debug, Deserialize)] @@ -28,38 +11,31 @@ struct GitHubUserInfo { id: u64, login: String, name: Option, + email: Option, avatar_url: Option, } impl SimpleOAuthProvider for GitHub { - fn get_authorize_url(&self) -> String { - String::from("https://github.com/login/oauth/authorize") - } - - fn get_token_url(&self) -> String { - String::from("https://github.com/login/oauth/access_token") - } - - fn get_scopes(&self) -> Vec { - vec!["user:read".into()] + fn authorize_url(&self) -> &str { + "https://github.com/login/oauth/authorize" } - fn get_user_info_url(&self) -> String { - String::from("https://api.github.com/user") + fn token_url(&self) -> &str { + "https://github.com/login/oauth/access_token" } - fn get_client_id(&self) -> String { - self.client_id.to_owned() + fn default_scopes(&self) -> Vec<&str> { + vec!["read:user"] } - fn get_client_secret(&self) -> String { - self.client_secret.to_owned() + fn user_info_url(&self) -> &str { + "https://api.github.com/user" } fn create_request_headers(&self) -> Vec<(String, String)> { vec![ ("Accept".into(), "application/vnd.github+json".into()), - ("User-Agent".into(), self.user_agent.clone()), + ("User-Agent".into(), "fa-sharp/simple-oauth".into()), ] } @@ -71,7 +47,9 @@ impl SimpleOAuthProvider for GitHub { Ok(UserInfo { id: info.id.to_string(), - name: info.name.unwrap_or(info.login), + name: info.name.or(Some(info.login)), + email: info.email, + email_verified: None, avatar_url: info.avatar_url, }) } diff --git a/server-new/crates/simple-oauth/src/common/google.rs b/server-new/crates/simple-oauth/src/common/google.rs index 7bbd93e..1a6d465 100644 --- a/server-new/crates/simple-oauth/src/common/google.rs +++ b/server-new/crates/simple-oauth/src/common/google.rs @@ -2,51 +2,35 @@ use serde::Deserialize; use crate::{SimpleOAuthProvider, types::UserInfo}; -pub struct Google { - client_id: String, - client_secret: String, -} - -impl Google { - pub fn new(client_id: impl Into, client_secret: impl Into) -> Self { - Self { - client_id: client_id.into(), - client_secret: client_secret.into(), - } - } -} +#[derive(Debug)] +pub struct Google; /// User info from Google API #[derive(Debug, Deserialize)] -pub struct GoogleUserInfo { +struct GoogleUserInfo { sub: String, - name: String, + name: Option, + preferred_username: Option, + email: Option, + email_verified: Option, picture: Option, } impl SimpleOAuthProvider for Google { - fn get_scopes(&self) -> Vec { - vec!["openid".into(), "profile".into()] - } - - fn get_authorize_url(&self) -> String { - "https://accounts.google.com/o/oauth2/v2/auth".into() - } - - fn get_token_url(&self) -> String { - "https://oauth2.googleapis.com/token".into() + fn default_scopes(&self) -> Vec<&str> { + vec!["openid", "profile"] } - fn get_user_info_url(&self) -> String { - "https://www.googleapis.com/oauth2/v3/userinfo".into() + fn authorize_url(&self) -> &str { + "https://accounts.google.com/o/oauth2/v2/auth" } - fn get_client_id(&self) -> String { - self.client_id.clone() + fn token_url(&self) -> &str { + "https://oauth2.googleapis.com/token" } - fn get_client_secret(&self) -> String { - self.client_secret.clone() + fn user_info_url(&self) -> &str { + "https://www.googleapis.com/oauth2/v3/userinfo" } fn extract_user_info(&self, val: serde_json::Value) -> Result { @@ -54,7 +38,9 @@ impl SimpleOAuthProvider for Google { Ok(UserInfo { id: user_info.sub, - name: user_info.name, + name: user_info.name.or(user_info.preferred_username), + email: user_info.email, + email_verified: user_info.email_verified, avatar_url: user_info.picture, }) } diff --git a/server-new/crates/simple-oauth/src/common/oidc.rs b/server-new/crates/simple-oauth/src/common/oidc.rs new file mode 100644 index 0000000..fac5da6 --- /dev/null +++ b/server-new/crates/simple-oauth/src/common/oidc.rs @@ -0,0 +1,85 @@ +use serde::Deserialize; + +use crate::{ + SimpleOAuthError, SimpleOAuthProvider, + types::{OidcDiscovery, UserInfo}, +}; + +#[derive(Debug)] +pub struct Oidc { + auth_endpoint: String, + token_endpoint: String, + userinfo_endpoint: String, +} + +/// Standard OIDC user info shape +#[derive(Debug, Deserialize)] +struct OidcUserInfo { + sub: String, + name: Option, + preferred_username: Option, + email: Option, + email_verified: Option, + picture: Option, +} + +impl Oidc { + pub fn from_config(config: OidcDiscovery) -> Self { + Self { + auth_endpoint: config.authorization_endpoint, + token_endpoint: config.token_endpoint, + userinfo_endpoint: config.userinfo_endpoint, + } + } + + /// Discover the OIDC config from the given URL. This will fail + /// if the discovery document is missing a token or userinfo endpoint. + pub async fn discover( + http_client: &reqwest::Client, + discovery_url: &str, + ) -> Result { + let discovery = http_client + .get(discovery_url) + .send() + .await? + .error_for_status()? + .json::() + .await?; + + Ok(Self { + auth_endpoint: discovery.authorization_endpoint, + token_endpoint: discovery.token_endpoint, + userinfo_endpoint: discovery.userinfo_endpoint, + }) + } +} + +impl SimpleOAuthProvider for Oidc { + fn authorize_url(&self) -> &str { + &self.auth_endpoint + } + + fn token_url(&self) -> &str { + &self.token_endpoint + } + + fn user_info_url(&self) -> &str { + &self.userinfo_endpoint + } + + fn default_scopes(&self) -> Vec<&str> { + vec!["openid", "profile"] + } + + fn extract_user_info(&self, val: serde_json::Value) -> Result { + let user_info: OidcUserInfo = serde_json::from_value(val)?; + + Ok(UserInfo { + id: user_info.sub, + name: user_info.name.or(user_info.preferred_username), + email: user_info.email, + email_verified: user_info.email_verified, + avatar_url: user_info.picture, + }) + } +} diff --git a/server-new/crates/simple-oauth/src/lib.rs b/server-new/crates/simple-oauth/src/lib.rs index 85af483..705313a 100644 --- a/server-new/crates/simple-oauth/src/lib.rs +++ b/server-new/crates/simple-oauth/src/lib.rs @@ -1,5 +1,6 @@ use oauth2::{ - CsrfToken, HttpClientError, RequestTokenError, TokenResponse, basic::BasicErrorResponse, + CsrfToken, HttpClientError, RequestTokenError, TokenResponse, + basic::{BasicClient, BasicErrorResponse}, }; pub mod common; @@ -8,7 +9,7 @@ pub mod types; pub use provider::SimpleOAuthProvider; -use crate::types::{AuthorizeUrl, StandardTokenResponse, UserInfo}; +use crate::types::{AuthorizeUrl, OAuthCredentials, StandardTokenResponse, UserInfo}; #[derive(Debug, thiserror::Error)] pub enum SimpleOAuthError { @@ -49,23 +50,25 @@ impl SimpleOAuthClient { pub fn authorize_url( &self, provider: &P, + credentials: OAuthCredentials<'_>, redirect_url: &str, scopes: Option>, ) -> Result { let oauth_client = - oauth2::basic::BasicClient::new(oauth2::ClientId::new(provider.get_client_id())) - .set_client_secret(oauth2::ClientSecret::new(provider.get_client_secret())) - .set_auth_uri(oauth2::AuthUrl::new(provider.get_authorize_url())?) + BasicClient::new(oauth2::ClientId::new(credentials.client_id.into_owned())) + .set_client_secret(oauth2::ClientSecret::new( + credentials.client_secret.into_owned(), + )) + .set_auth_uri(oauth2::AuthUrl::new(provider.authorize_url().into())?) .set_redirect_uri(oauth2::RedirectUrl::new(redirect_url.into())?); let (pkce_challenge, pkce_verifier) = oauth2::PkceCodeChallenge::new_random_sha256(); let (url, state) = oauth_client .authorize_url(CsrfToken::new_random) .add_scopes( scopes - .map(|scopes| scopes.into_iter().map(|s| s.to_owned()).collect()) - .unwrap_or_else(|| provider.get_scopes()) + .unwrap_or_else(|| provider.default_scopes()) .into_iter() - .map(oauth2::Scope::new), + .map(|s| oauth2::Scope::new(s.into())), ) .set_pkce_challenge(pkce_challenge) .url(); @@ -80,15 +83,18 @@ impl SimpleOAuthClient { pub async fn exchange_code( &self, provider: &P, + credentials: OAuthCredentials<'_>, redirect_url: &str, code: &str, pkce_verifier: Option<&str>, ) -> Result { let oauth_client = - oauth2::basic::BasicClient::new(oauth2::ClientId::new(provider.get_client_id())) - .set_client_secret(oauth2::ClientSecret::new(provider.get_client_secret())) + BasicClient::new(oauth2::ClientId::new(credentials.client_id.into_owned())) + .set_client_secret(oauth2::ClientSecret::new( + credentials.client_secret.into_owned(), + )) .set_redirect_uri(oauth2::RedirectUrl::new(redirect_url.into())?) - .set_token_uri(oauth2::TokenUrl::new(provider.get_token_url())?); + .set_token_uri(oauth2::TokenUrl::new(provider.token_url().into())?); let mut token_request = oauth_client.exchange_code(oauth2::AuthorizationCode::new(code.into())); if let Some(verifier) = pkce_verifier { @@ -111,7 +117,7 @@ impl SimpleOAuthClient { ) -> Result { let mut user_info_request = self .http_client - .get(provider.get_user_info_url()) + .get(provider.user_info_url()) .bearer_auth(access_token); for (name, val) in provider.create_request_headers() { user_info_request = user_info_request.header(name, val); diff --git a/server-new/crates/simple-oauth/src/provider.rs b/server-new/crates/simple-oauth/src/provider.rs index a33b900..fd2be46 100644 --- a/server-new/crates/simple-oauth/src/provider.rs +++ b/server-new/crates/simple-oauth/src/provider.rs @@ -1,18 +1,15 @@ +use std::fmt::Debug; + use crate::types::UserInfo; /// Trait for all OAuth providers -pub trait SimpleOAuthProvider: Send + Sync { - fn get_scopes(&self) -> Vec; - fn get_authorize_url(&self) -> String; - fn get_token_url(&self) -> String; - fn get_user_info_url(&self) -> String; - fn get_client_id(&self) -> String; - fn get_client_secret(&self) -> String; +pub trait SimpleOAuthProvider: Debug + Send + Sync { + fn authorize_url(&self) -> &str; + fn token_url(&self) -> &str; + fn user_info_url(&self) -> &str; + fn default_scopes(&self) -> Vec<&str>; fn create_request_headers(&self) -> Vec<(String, String)> { vec![] } - fn extract_user_info( - &self, - user_info: serde_json::Value, - ) -> Result; + fn extract_user_info(&self, val: serde_json::Value) -> Result; } diff --git a/server-new/crates/simple-oauth/src/types.rs b/server-new/crates/simple-oauth/src/types.rs index 89d4d8f..8b163ee 100644 --- a/server-new/crates/simple-oauth/src/types.rs +++ b/server-new/crates/simple-oauth/src/types.rs @@ -1,17 +1,83 @@ +use std::{borrow::Cow, fmt::Debug}; + +use serde::Deserialize; + +/// OAuth2 authorization redirect URL, along with the state and PKCE verifier pub struct AuthorizeUrl { pub url: oauth2::url::Url, pub state: String, pub pkce_verifier: String, } +impl Debug for AuthorizeUrl { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("AuthorizeUrl") + .field("url", &"--redacted--") + .field("state", &self.state) + .field("pkce_verifier", &"--redacted--") + .finish() + } +} + +/// User info returned by the OAuth provider +#[derive(Debug)] +pub struct UserInfo { + /// The ID of the user at the OAuth provider + pub id: String, + /// The user's display name + pub name: Option, + /// The user's email. Likely will not be included unless you add the proper email scope for the provider. + /// + /// ⚠️ Do not rely on this for identifying the user. Use the `id` and the provider name. + pub email: Option, + /// Whether the user's email is verified. Not all providers return this in the user info. + pub email_verified: Option, + /// The URL of the user's picture/avatar + pub avatar_url: Option, +} +/// Standard OAuth2 token response pub struct StandardTokenResponse { pub access_token: String, pub refresh_token: Option, pub expires_in: Option, } +impl Debug for StandardTokenResponse { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("StandardTokenResponse") + .field("access_token", &"--redacted--") + .field("refresh_token", &"--redacted--") + .field("expires_in", &self.expires_in) + .finish() + } +} -pub struct UserInfo { - pub id: String, - pub name: String, - pub avatar_url: Option, +pub struct OAuthCredentials<'a> { + pub client_id: Cow<'a, str>, + pub client_secret: Cow<'a, str>, +} +impl<'a> OAuthCredentials<'a> { + pub fn new(client_id: impl Into>, client_secret: impl Into>) -> Self { + Self { + client_id: client_id.into(), + client_secret: client_secret.into(), + } + } +} +impl<'a> Debug for OAuthCredentials<'a> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("OAuthCredentials") + .field("client_id", &self.client_id) + .field("client_secret", &"--redacted--") + .finish() + } +} + +/// OIDC discovery document +#[derive(Debug, Default, Deserialize)] +pub struct OidcDiscovery { + pub issuer: String, + pub authorization_endpoint: String, + pub token_endpoint: String, + pub userinfo_endpoint: String, + pub scopes_supported: Option>, } diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index 927f152..30787bd 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -1,8 +1,8 @@ use futures::future::BoxFuture; use serde::{Deserialize, Serialize}; use simple_oauth::{ - SimpleOAuthProvider, - types::{AuthorizeUrl, StandardTokenResponse, UserInfo}, + SimpleOAuthClient, SimpleOAuthProvider, + types::{AuthorizeUrl, OAuthCredentials, StandardTokenResponse, UserInfo}, }; use subtle::ConstantTimeEq; use tower_sessions::Session; @@ -45,6 +45,7 @@ impl OAuthProviderEnum { /// Trait for all OAuth providers pub trait OAuthProvider: Send + Sync { fn get_inner_provider(&self) -> Box; + fn get_credentials(&self) -> OAuthCredentials<'_>; fn find_linked_user<'a>( &self, db: &'a mut DbService, @@ -85,14 +86,14 @@ impl<'a> OAuthService<'a> { session: &Session, ) -> AuthResult { let provider = self.get_provider(provider)?; - let client = simple_oauth::SimpleOAuthClient::with_http_client(self.http_client.clone()); - + let client = SimpleOAuthClient::with_http_client(self.http_client.clone()); let AuthorizeUrl { url, state, pkce_verifier, } = client.authorize_url( provider.get_inner_provider().as_ref(), + provider.get_credentials(), &self.get_redirect_url(callback_path), None, )?; @@ -127,10 +128,11 @@ impl<'a> OAuthService<'a> { // Exchange code for token let provider = self.get_provider(provider)?; - let client = simple_oauth::SimpleOAuthClient::with_http_client(self.http_client.clone()); + let client = SimpleOAuthClient::with_http_client(self.http_client.clone()); let response = client .exchange_code( provider.get_inner_provider().as_ref(), + provider.get_credentials(), &self.get_redirect_url(callback_path), code, Some(&pkce_verifier), @@ -147,7 +149,7 @@ impl<'a> OAuthService<'a> { active_session: Option, ) -> AuthResult { let provider = self.get_provider(provider)?; - let client = simple_oauth::SimpleOAuthClient::with_http_client(self.http_client.clone()); + let client = SimpleOAuthClient::with_http_client(self.http_client.clone()); // Get user info from provider let user_info = client diff --git a/server-new/src/services/auth/oauth/discord.rs b/server-new/src/services/auth/oauth/discord.rs index d593464..f0875cd 100644 --- a/server-new/src/services/auth/oauth/discord.rs +++ b/server-new/src/services/auth/oauth/discord.rs @@ -1,6 +1,6 @@ use futures::future::BoxFuture; use serde::Deserialize; -use simple_oauth::{common::discord::Discord, types::UserInfo}; +use simple_oauth::types::{OAuthCredentials, UserInfo}; use crate::{ db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, @@ -27,10 +27,14 @@ impl DiscordOAuthProvider { impl OAuthProvider for DiscordOAuthProvider { fn get_inner_provider(&self) -> Box { - Box::new(Discord::new( - self.config.client_id, + Box::new(simple_oauth::common::discord::Discord) + } + + fn get_credentials(&self) -> OAuthCredentials<'_> { + OAuthCredentials::new( + self.config.client_id.to_string(), &self.config.client_secret, - )) + ) } fn find_linked_user<'a>( @@ -58,7 +62,7 @@ impl OAuthProvider for DiscordOAuthProvider { fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { discord_id: Some(&user_data.id), - name: &user_data.name, + name: &user_data.name.as_deref().unwrap_or_default(), avatar_url: user_data.avatar_url.as_deref(), ..Default::default() } diff --git a/server-new/src/services/auth/oauth/github.rs b/server-new/src/services/auth/oauth/github.rs index 15d976b..060379a 100644 --- a/server-new/src/services/auth/oauth/github.rs +++ b/server-new/src/services/auth/oauth/github.rs @@ -1,6 +1,9 @@ use futures::future::BoxFuture; use serde::Deserialize; -use simple_oauth::{SimpleOAuthProvider, common::github::GitHub, types::UserInfo}; +use simple_oauth::{ + SimpleOAuthProvider, + types::{OAuthCredentials, UserInfo}, +}; use crate::{ db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, @@ -27,11 +30,11 @@ impl GitHubOAuthProvider { impl OAuthProvider for GitHubOAuthProvider { fn get_inner_provider(&self) -> Box { - Box::new(GitHub::new( - &self.config.client_id, - &self.config.client_secret, - "fa-sharp/rs-chat", - )) + Box::new(simple_oauth::common::github::GitHub) + } + + fn get_credentials(&self) -> OAuthCredentials<'_> { + OAuthCredentials::new(&self.config.client_id, &self.config.client_secret) } fn find_linked_user<'a>( @@ -59,7 +62,7 @@ impl OAuthProvider for GitHubOAuthProvider { fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { github_id: Some(&user_data.id), - name: &user_data.name, + name: &user_data.name.as_deref().unwrap_or_default(), avatar_url: user_data.avatar_url.as_deref(), ..Default::default() } diff --git a/server-new/src/services/auth/oauth/google.rs b/server-new/src/services/auth/oauth/google.rs index f77b043..a48901d 100644 --- a/server-new/src/services/auth/oauth/google.rs +++ b/server-new/src/services/auth/oauth/google.rs @@ -1,6 +1,6 @@ use futures::future::BoxFuture; use serde::Deserialize; -use simple_oauth::{common::google::Google, types::UserInfo}; +use simple_oauth::types::{OAuthCredentials, UserInfo}; use crate::{ db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, @@ -27,10 +27,11 @@ impl GoogleOAuthProvider { impl OAuthProvider for GoogleOAuthProvider { fn get_inner_provider(&self) -> Box { - Box::new(Google::new( - &self.config.client_id, - &self.config.client_secret, - )) + Box::new(simple_oauth::common::google::Google) + } + + fn get_credentials(&self) -> OAuthCredentials<'_> { + OAuthCredentials::new(&self.config.client_id, &self.config.client_secret) } fn find_linked_user<'a>( @@ -58,7 +59,7 @@ impl OAuthProvider for GoogleOAuthProvider { fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { google_id: Some(&user_data.id), - name: &user_data.name, + name: &user_data.name.as_deref().unwrap_or_default(), avatar_url: user_data.avatar_url.as_deref(), ..Default::default() } From 002617aa24742c64372b0094ed700e0107a4d79e Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sun, 28 Jun 2026 00:44:33 -0400 Subject: [PATCH 041/111] moar oauth tweaks --- server-new/crates/simple-oauth/Cargo.toml | 2 +- .../crates/simple-oauth/src/common/discord.rs | 4 ++-- server-new/crates/simple-oauth/src/common/github.rs | 4 ++-- server-new/crates/simple-oauth/src/common/google.rs | 4 ++-- server-new/crates/simple-oauth/src/common/oidc.rs | 4 ++-- server-new/crates/simple-oauth/src/lib.rs | 8 ++++---- server-new/crates/simple-oauth/src/provider.rs | 2 +- server-new/src/api/auth.rs | 2 +- server-new/src/plugins/database.rs | 13 ++++++++++--- server-new/src/plugins/redis.rs | 6 +++--- server-new/src/plugins/session.rs | 11 ++++------- 11 files changed, 32 insertions(+), 28 deletions(-) diff --git a/server-new/crates/simple-oauth/Cargo.toml b/server-new/crates/simple-oauth/Cargo.toml index 99c721b..328c433 100644 --- a/server-new/crates/simple-oauth/Cargo.toml +++ b/server-new/crates/simple-oauth/Cargo.toml @@ -2,7 +2,7 @@ name = "simple-oauth" version = "0.1.0" edition = "2024" -description = "Simple OAuth2 authorization crate" +description = "Simple OAuth2 login and authorization" [dependencies] oauth2 = { version = "5", default-features = false } diff --git a/server-new/crates/simple-oauth/src/common/discord.rs b/server-new/crates/simple-oauth/src/common/discord.rs index 119a1cb..6a65b3b 100644 --- a/server-new/crates/simple-oauth/src/common/discord.rs +++ b/server-new/crates/simple-oauth/src/common/discord.rs @@ -25,8 +25,8 @@ impl SimpleOAuthProvider for Discord { "https://discord.com/api/oauth2/token" } - fn default_scopes(&self) -> Vec<&str> { - vec!["identify"] + fn default_scopes(&self) -> &'static [&'static str] { + &["identify"] } fn user_info_url(&self) -> &str { diff --git a/server-new/crates/simple-oauth/src/common/github.rs b/server-new/crates/simple-oauth/src/common/github.rs index a5b1119..8fab86c 100644 --- a/server-new/crates/simple-oauth/src/common/github.rs +++ b/server-new/crates/simple-oauth/src/common/github.rs @@ -24,8 +24,8 @@ impl SimpleOAuthProvider for GitHub { "https://github.com/login/oauth/access_token" } - fn default_scopes(&self) -> Vec<&str> { - vec!["read:user"] + fn default_scopes(&self) -> &'static [&'static str] { + &["read:user"] } fn user_info_url(&self) -> &str { diff --git a/server-new/crates/simple-oauth/src/common/google.rs b/server-new/crates/simple-oauth/src/common/google.rs index 1a6d465..e93b13d 100644 --- a/server-new/crates/simple-oauth/src/common/google.rs +++ b/server-new/crates/simple-oauth/src/common/google.rs @@ -17,8 +17,8 @@ struct GoogleUserInfo { } impl SimpleOAuthProvider for Google { - fn default_scopes(&self) -> Vec<&str> { - vec!["openid", "profile"] + fn default_scopes(&self) -> &'static [&'static str] { + &["openid", "profile"] } fn authorize_url(&self) -> &str { diff --git a/server-new/crates/simple-oauth/src/common/oidc.rs b/server-new/crates/simple-oauth/src/common/oidc.rs index fac5da6..4a1d2ee 100644 --- a/server-new/crates/simple-oauth/src/common/oidc.rs +++ b/server-new/crates/simple-oauth/src/common/oidc.rs @@ -67,8 +67,8 @@ impl SimpleOAuthProvider for Oidc { &self.userinfo_endpoint } - fn default_scopes(&self) -> Vec<&str> { - vec!["openid", "profile"] + fn default_scopes(&self) -> &'static [&'static str] { + &["openid", "profile"] } fn extract_user_info(&self, val: serde_json::Value) -> Result { diff --git a/server-new/crates/simple-oauth/src/lib.rs b/server-new/crates/simple-oauth/src/lib.rs index 705313a..af9b5f8 100644 --- a/server-new/crates/simple-oauth/src/lib.rs +++ b/server-new/crates/simple-oauth/src/lib.rs @@ -52,7 +52,7 @@ impl SimpleOAuthClient { provider: &P, credentials: OAuthCredentials<'_>, redirect_url: &str, - scopes: Option>, + custom_scopes: Option<&[&str]>, ) -> Result { let oauth_client = BasicClient::new(oauth2::ClientId::new(credentials.client_id.into_owned())) @@ -65,10 +65,10 @@ impl SimpleOAuthClient { let (url, state) = oauth_client .authorize_url(CsrfToken::new_random) .add_scopes( - scopes - .unwrap_or_else(|| provider.default_scopes()) + custom_scopes + .unwrap_or(provider.default_scopes()) .into_iter() - .map(|s| oauth2::Scope::new(s.into())), + .map(|s| oauth2::Scope::new((*s).to_owned())), ) .set_pkce_challenge(pkce_challenge) .url(); diff --git a/server-new/crates/simple-oauth/src/provider.rs b/server-new/crates/simple-oauth/src/provider.rs index fd2be46..c764703 100644 --- a/server-new/crates/simple-oauth/src/provider.rs +++ b/server-new/crates/simple-oauth/src/provider.rs @@ -7,7 +7,7 @@ pub trait SimpleOAuthProvider: Debug + Send + Sync { fn authorize_url(&self) -> &str; fn token_url(&self) -> &str; fn user_info_url(&self) -> &str; - fn default_scopes(&self) -> Vec<&str>; + fn default_scopes(&self) -> &'static [&'static str]; fn create_request_headers(&self) -> Vec<(String, String)> { vec![] } diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index 5ca3734..c881f8d 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -19,7 +19,7 @@ pub fn routes() -> axum::Router { .route("/login/{provider}", routing::get(login_handler)) .route("/login/{provider}/callback", routing::get(callback_handler)) .route("/user", routing::get(get_user_handler)) - .route("/logout", routing::post(logout_handler)) + .route("/logout", routing::get(logout_handler).post(logout_handler)) } fn callback_path(route_prefix: &'static str, provider: OAuthProviderEnum) -> String { diff --git a/server-new/src/plugins/database.rs b/server-new/src/plugins/database.rs index d1fe098..e98ef47 100644 --- a/server-new/src/plugins/database.rs +++ b/server-new/src/plugins/database.rs @@ -2,7 +2,7 @@ use anyhow::Context; use axum_plugin::AdHocPlugin; use diesel_async::{ AsyncMigrationHarness, AsyncPgConnection, - pooled_connection::{AsyncDieselConnectionManager, deadpool::Pool}, + pooled_connection::{AsyncDieselConnectionManager, ManagerConfig, deadpool::Pool}, }; use diesel_migrations::{EmbeddedMigrations, MigrationHarness}; @@ -15,8 +15,15 @@ pub fn plugin() -> AdHocPlugin { .on_init(async |mut state| { let app_config = state.get::().context("missing config")?; - let manager = - AsyncDieselConnectionManager::::new(&app_config.database.url); + let manager = AsyncDieselConnectionManager::::new_with_config( + &app_config.database.url, + { + let mut config = ManagerConfig::default(); + config.recycling_method = + diesel_async::pooled_connection::RecyclingMethod::Fast; + config + }, + ); let pool: DbPool = Pool::builder(manager).build()?; let cxn = pool.get().await.context("failed to connect to database")?; diff --git a/server-new/src/plugins/redis.rs b/server-new/src/plugins/redis.rs index ce4d1c1..ec334b5 100644 --- a/server-new/src/plugins/redis.rs +++ b/server-new/src/plugins/redis.rs @@ -4,15 +4,15 @@ use anyhow::Context; use axum_plugin::AdHocPlugin; use fred::prelude::*; -use crate::{config::AppConfig, db::DbPool, state::AppState}; +use crate::{config::AppConfig, state::AppState}; const DEFAULT_TIMEOUT: Duration = Duration::from_secs(8); +const POOL_SIZE: usize = 4; pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Redis") .on_init(async |mut state| { let app_config = state.get::().context("no config")?; - let db_pool = state.get::().context("no database pool")?; let config = Config::from_url(&app_config.redis.url).context("invalid Redis URL")?; let pool = Builder::from_config(config) .with_connection_config(|c| { @@ -23,7 +23,7 @@ pub fn plugin() -> AdHocPlugin { .with_performance_config(|c| { c.default_command_timeout = DEFAULT_TIMEOUT; }) - .build_pool(db_pool.status().max_size)?; // same size as database pool + .build_pool(POOL_SIZE)?; pool.init().await.context("failed to connect to Redis")?; tracing::info!("Connected to Redis"); diff --git a/server-new/src/plugins/session.rs b/server-new/src/plugins/session.rs index 6f111b5..5a49277 100644 --- a/server-new/src/plugins/session.rs +++ b/server-new/src/plugins/session.rs @@ -2,10 +2,7 @@ use std::time::Duration; use anyhow::{Context, bail}; use axum_plugin::AdHocPlugin; -use tower_sessions::{ - CachingSessionStore, Expiry, SessionManagerLayer, - cookie::{Key, SameSite}, -}; +use tower_sessions::{CachingSessionStore, Expiry, SessionManagerLayer, cookie}; use tower_sessions_redis_store::RedisStore; use crate::{ @@ -54,12 +51,12 @@ pub fn plugin() -> AdHocPlugin { // Add session / cookie management to router let session_layer = SessionManagerLayer::new(session_store) .with_name(state.config.auth.cookie_name.clone()) - .with_expiry(Expiry::OnSessionEnd) - .with_private(Key::derive_from(&cookie_key)) + .with_expiry(Expiry::OnInactivity(cookie::time::Duration::minutes(15))) // default short session for login/OAuth + .with_private(cookie::Key::derive_from(&cookie_key)) .with_path("/") .with_secure(true) .with_http_only(true) - .with_same_site(SameSite::Lax); + .with_same_site(cookie::SameSite::Lax); Ok(router.layer(session_layer)) }) From 398458ba31de1c878ecdfadb84b302b37dad1b1d Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sun, 28 Jun 2026 16:45:47 -0400 Subject: [PATCH 042/111] use extracted simple-oauth crate --- server-new/Cargo.lock | 39 ++++- server-new/Cargo.toml | 7 +- server-new/crates/simple-oauth/Cargo.toml | 18 --- server-new/crates/simple-oauth/src/common.rs | 5 - .../crates/simple-oauth/src/common/discord.rs | 53 ------- .../crates/simple-oauth/src/common/github.rs | 56 -------- .../crates/simple-oauth/src/common/google.rs | 47 ------ .../crates/simple-oauth/src/common/oidc.rs | 85 ----------- server-new/crates/simple-oauth/src/lib.rs | 136 ------------------ .../crates/simple-oauth/src/provider.rs | 15 -- server-new/crates/simple-oauth/src/types.rs | 83 ----------- server-new/src/services/auth/oauth.rs | 82 ++++++----- server-new/src/services/auth/oauth/discord.rs | 10 +- server-new/src/services/auth/oauth/github.rs | 10 +- server-new/src/services/auth/oauth/google.rs | 10 +- 15 files changed, 104 insertions(+), 552 deletions(-) delete mode 100644 server-new/crates/simple-oauth/Cargo.toml delete mode 100644 server-new/crates/simple-oauth/src/common.rs delete mode 100644 server-new/crates/simple-oauth/src/common/discord.rs delete mode 100644 server-new/crates/simple-oauth/src/common/github.rs delete mode 100644 server-new/crates/simple-oauth/src/common/google.rs delete mode 100644 server-new/crates/simple-oauth/src/common/oidc.rs delete mode 100644 server-new/crates/simple-oauth/src/lib.rs delete mode 100644 server-new/crates/simple-oauth/src/provider.rs delete mode 100644 server-new/crates/simple-oauth/src/types.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index aa6a989..e8389ff 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -232,6 +232,31 @@ dependencies = [ "generic-array", ] +[[package]] +name = "bon" +version = "3.9.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a602c73c7b0148ec6d12af6fd5cc7a46e2eacc8878271a999abac56eed12f561" +dependencies = [ + "bon-macros", + "rustversion", +] + +[[package]] +name = "bon-macros" +version = "3.9.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6dee98b0db6a962de883bf5d20362dee4d7ca0d12fe39a7c6c73c844e1cd7c1f" +dependencies = [ + "darling 0.21.3", + "ident_case", + "prettyplease", + "proc-macro2", + "quote", + "rustversion", + "syn", +] + [[package]] name = "bumpalo" version = "3.20.3" @@ -1673,6 +1698,16 @@ dependencies = [ "zerocopy", ] +[[package]] +name = "prettyplease" +version = "0.2.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" +dependencies = [ + "proc-macro2", + "syn", +] + [[package]] name = "proc-macro2" version = "1.0.106" @@ -1965,7 +2000,6 @@ dependencies = [ "serde_json", "serde_with", "simple-oauth", - "subtle", "thiserror 2.0.18", "tokio", "tower", @@ -2294,12 +2328,15 @@ checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" [[package]] name = "simple-oauth" version = "0.1.0" +source = "git+https://github.com/fa-sharp/simple-oauth-rs?rev=45ab590#45ab590710c4830d509c8b5b1781fd19d381525b" dependencies = [ + "bon", "oauth2", "oauth2-reqwest", "reqwest", "serde", "serde_json", + "subtle", "thiserror 2.0.18", ] diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index dc876a7..8d46425 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -51,8 +51,11 @@ serde_with = { default-features = false, features = ["macros"] } -simple-oauth = { path = "crates/simple-oauth/", features = ["default-tls"] } -subtle = "2.6.1" +simple-oauth = { + git = "https://github.com/fa-sharp/simple-oauth-rs", + rev = "45ab590", + features = ["default-tls"] +} thiserror = "2.0.18" tokio = { version = "1.52.3", diff --git a/server-new/crates/simple-oauth/Cargo.toml b/server-new/crates/simple-oauth/Cargo.toml deleted file mode 100644 index 328c433..0000000 --- a/server-new/crates/simple-oauth/Cargo.toml +++ /dev/null @@ -1,18 +0,0 @@ -[package] -name = "simple-oauth" -version = "0.1.0" -edition = "2024" -description = "Simple OAuth2 login and authorization" - -[dependencies] -oauth2 = { version = "5", default-features = false } -oauth2-reqwest = "0.1.0-alpha.3" -reqwest = { version = "0.13", default-features = false, features = ["json"] } -serde = "1.0" -serde_json = "1.0" -thiserror = "2.0" - -[features] -default = ["default-tls"] -default-tls = ["reqwest/default-tls"] -native-tls = ["reqwest/native-tls"] diff --git a/server-new/crates/simple-oauth/src/common.rs b/server-new/crates/simple-oauth/src/common.rs deleted file mode 100644 index 844d508..0000000 --- a/server-new/crates/simple-oauth/src/common.rs +++ /dev/null @@ -1,5 +0,0 @@ -pub mod discord; -pub mod github; -pub mod google; - -pub mod oidc; diff --git a/server-new/crates/simple-oauth/src/common/discord.rs b/server-new/crates/simple-oauth/src/common/discord.rs deleted file mode 100644 index 6a65b3b..0000000 --- a/server-new/crates/simple-oauth/src/common/discord.rs +++ /dev/null @@ -1,53 +0,0 @@ -use serde::Deserialize; - -use crate::{SimpleOAuthProvider, types::UserInfo}; - -#[derive(Debug)] -pub struct Discord; - -/// User info returned from Discord API -#[derive(Debug, Deserialize)] -struct DiscordUserInfo { - id: String, - username: String, - global_name: Option, - email: Option, - verified: Option, - avatar: Option, -} - -impl SimpleOAuthProvider for Discord { - fn authorize_url(&self) -> &str { - "https://discord.com/oauth2/authorize" - } - - fn token_url(&self) -> &str { - "https://discord.com/api/oauth2/token" - } - - fn default_scopes(&self) -> &'static [&'static str] { - &["identify"] - } - - fn user_info_url(&self) -> &str { - "https://discord.com/api/v9/users/@me" - } - - fn extract_user_info(&self, val: serde_json::Value) -> Result { - let user_info: DiscordUserInfo = serde_json::from_value(val)?; - let avatar_url = user_info.avatar.as_ref().map(|avatar| { - format!( - "https://cdn.discordapp.com/avatars/{}/{}.png", - user_info.id, avatar - ) - }); - - Ok(UserInfo { - id: user_info.id, - email: user_info.email, - email_verified: user_info.verified, - name: user_info.global_name.or(Some(user_info.username)), - avatar_url, - }) - } -} diff --git a/server-new/crates/simple-oauth/src/common/github.rs b/server-new/crates/simple-oauth/src/common/github.rs deleted file mode 100644 index 8fab86c..0000000 --- a/server-new/crates/simple-oauth/src/common/github.rs +++ /dev/null @@ -1,56 +0,0 @@ -use serde::Deserialize; - -use crate::{SimpleOAuthProvider, types::UserInfo}; - -#[derive(Debug)] -pub struct GitHub; - -/// User info returned from GitHub API -#[derive(Debug, Deserialize)] -struct GitHubUserInfo { - id: u64, - login: String, - name: Option, - email: Option, - avatar_url: Option, -} - -impl SimpleOAuthProvider for GitHub { - fn authorize_url(&self) -> &str { - "https://github.com/login/oauth/authorize" - } - - fn token_url(&self) -> &str { - "https://github.com/login/oauth/access_token" - } - - fn default_scopes(&self) -> &'static [&'static str] { - &["read:user"] - } - - fn user_info_url(&self) -> &str { - "https://api.github.com/user" - } - - fn create_request_headers(&self) -> Vec<(String, String)> { - vec![ - ("Accept".into(), "application/vnd.github+json".into()), - ("User-Agent".into(), "fa-sharp/simple-oauth".into()), - ] - } - - fn extract_user_info( - &self, - user_info: serde_json::Value, - ) -> Result { - let info: GitHubUserInfo = serde_json::from_value(user_info)?; - - Ok(UserInfo { - id: info.id.to_string(), - name: info.name.or(Some(info.login)), - email: info.email, - email_verified: None, - avatar_url: info.avatar_url, - }) - } -} diff --git a/server-new/crates/simple-oauth/src/common/google.rs b/server-new/crates/simple-oauth/src/common/google.rs deleted file mode 100644 index e93b13d..0000000 --- a/server-new/crates/simple-oauth/src/common/google.rs +++ /dev/null @@ -1,47 +0,0 @@ -use serde::Deserialize; - -use crate::{SimpleOAuthProvider, types::UserInfo}; - -#[derive(Debug)] -pub struct Google; - -/// User info from Google API -#[derive(Debug, Deserialize)] -struct GoogleUserInfo { - sub: String, - name: Option, - preferred_username: Option, - email: Option, - email_verified: Option, - picture: Option, -} - -impl SimpleOAuthProvider for Google { - fn default_scopes(&self) -> &'static [&'static str] { - &["openid", "profile"] - } - - fn authorize_url(&self) -> &str { - "https://accounts.google.com/o/oauth2/v2/auth" - } - - fn token_url(&self) -> &str { - "https://oauth2.googleapis.com/token" - } - - fn user_info_url(&self) -> &str { - "https://www.googleapis.com/oauth2/v3/userinfo" - } - - fn extract_user_info(&self, val: serde_json::Value) -> Result { - let user_info: GoogleUserInfo = serde_json::from_value(val)?; - - Ok(UserInfo { - id: user_info.sub, - name: user_info.name.or(user_info.preferred_username), - email: user_info.email, - email_verified: user_info.email_verified, - avatar_url: user_info.picture, - }) - } -} diff --git a/server-new/crates/simple-oauth/src/common/oidc.rs b/server-new/crates/simple-oauth/src/common/oidc.rs deleted file mode 100644 index 4a1d2ee..0000000 --- a/server-new/crates/simple-oauth/src/common/oidc.rs +++ /dev/null @@ -1,85 +0,0 @@ -use serde::Deserialize; - -use crate::{ - SimpleOAuthError, SimpleOAuthProvider, - types::{OidcDiscovery, UserInfo}, -}; - -#[derive(Debug)] -pub struct Oidc { - auth_endpoint: String, - token_endpoint: String, - userinfo_endpoint: String, -} - -/// Standard OIDC user info shape -#[derive(Debug, Deserialize)] -struct OidcUserInfo { - sub: String, - name: Option, - preferred_username: Option, - email: Option, - email_verified: Option, - picture: Option, -} - -impl Oidc { - pub fn from_config(config: OidcDiscovery) -> Self { - Self { - auth_endpoint: config.authorization_endpoint, - token_endpoint: config.token_endpoint, - userinfo_endpoint: config.userinfo_endpoint, - } - } - - /// Discover the OIDC config from the given URL. This will fail - /// if the discovery document is missing a token or userinfo endpoint. - pub async fn discover( - http_client: &reqwest::Client, - discovery_url: &str, - ) -> Result { - let discovery = http_client - .get(discovery_url) - .send() - .await? - .error_for_status()? - .json::() - .await?; - - Ok(Self { - auth_endpoint: discovery.authorization_endpoint, - token_endpoint: discovery.token_endpoint, - userinfo_endpoint: discovery.userinfo_endpoint, - }) - } -} - -impl SimpleOAuthProvider for Oidc { - fn authorize_url(&self) -> &str { - &self.auth_endpoint - } - - fn token_url(&self) -> &str { - &self.token_endpoint - } - - fn user_info_url(&self) -> &str { - &self.userinfo_endpoint - } - - fn default_scopes(&self) -> &'static [&'static str] { - &["openid", "profile"] - } - - fn extract_user_info(&self, val: serde_json::Value) -> Result { - let user_info: OidcUserInfo = serde_json::from_value(val)?; - - Ok(UserInfo { - id: user_info.sub, - name: user_info.name.or(user_info.preferred_username), - email: user_info.email, - email_verified: user_info.email_verified, - avatar_url: user_info.picture, - }) - } -} diff --git a/server-new/crates/simple-oauth/src/lib.rs b/server-new/crates/simple-oauth/src/lib.rs deleted file mode 100644 index af9b5f8..0000000 --- a/server-new/crates/simple-oauth/src/lib.rs +++ /dev/null @@ -1,136 +0,0 @@ -use oauth2::{ - CsrfToken, HttpClientError, RequestTokenError, TokenResponse, - basic::{BasicClient, BasicErrorResponse}, -}; - -pub mod common; -mod provider; -pub mod types; - -pub use provider::SimpleOAuthProvider; - -use crate::types::{AuthorizeUrl, OAuthCredentials, StandardTokenResponse, UserInfo}; - -#[derive(Debug, thiserror::Error)] -pub enum SimpleOAuthError { - #[error(transparent)] - Request(#[from] reqwest::Error), - #[error("invalid url: {0}")] - ParseUrl(#[from] oauth2::url::ParseError), - #[error("token exchange error: {0}")] - TokenExchange(#[from] RequestTokenError, BasicErrorResponse>), - #[error("deserialization error: {0}")] - Deserialization(#[from] serde_json::Error), -} - -pub struct SimpleOAuthClient { - http_client: reqwest::Client, - oauth_client: oauth2_reqwest::ReqwestClient, -} - -impl SimpleOAuthClient { - pub fn new() -> Result { - let http_client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build()?; - - Ok(Self { - oauth_client: oauth2_reqwest::ReqwestClient::from(http_client.clone()), - http_client, - }) - } - - pub fn with_http_client(http_client: reqwest::Client) -> Self { - Self { - oauth_client: oauth2_reqwest::ReqwestClient::from(http_client.clone()), - http_client, - } - } - - pub fn authorize_url( - &self, - provider: &P, - credentials: OAuthCredentials<'_>, - redirect_url: &str, - custom_scopes: Option<&[&str]>, - ) -> Result { - let oauth_client = - BasicClient::new(oauth2::ClientId::new(credentials.client_id.into_owned())) - .set_client_secret(oauth2::ClientSecret::new( - credentials.client_secret.into_owned(), - )) - .set_auth_uri(oauth2::AuthUrl::new(provider.authorize_url().into())?) - .set_redirect_uri(oauth2::RedirectUrl::new(redirect_url.into())?); - let (pkce_challenge, pkce_verifier) = oauth2::PkceCodeChallenge::new_random_sha256(); - let (url, state) = oauth_client - .authorize_url(CsrfToken::new_random) - .add_scopes( - custom_scopes - .unwrap_or(provider.default_scopes()) - .into_iter() - .map(|s| oauth2::Scope::new((*s).to_owned())), - ) - .set_pkce_challenge(pkce_challenge) - .url(); - - Ok(AuthorizeUrl { - url, - state: state.into_secret(), - pkce_verifier: pkce_verifier.into_secret(), - }) - } - - pub async fn exchange_code( - &self, - provider: &P, - credentials: OAuthCredentials<'_>, - redirect_url: &str, - code: &str, - pkce_verifier: Option<&str>, - ) -> Result { - let oauth_client = - BasicClient::new(oauth2::ClientId::new(credentials.client_id.into_owned())) - .set_client_secret(oauth2::ClientSecret::new( - credentials.client_secret.into_owned(), - )) - .set_redirect_uri(oauth2::RedirectUrl::new(redirect_url.into())?) - .set_token_uri(oauth2::TokenUrl::new(provider.token_url().into())?); - let mut token_request = - oauth_client.exchange_code(oauth2::AuthorizationCode::new(code.into())); - if let Some(verifier) = pkce_verifier { - token_request = - token_request.set_pkce_verifier(oauth2::PkceCodeVerifier::new(verifier.into())); - } - let token = token_request.request_async(&self.oauth_client).await?; - - Ok(StandardTokenResponse { - access_token: token.access_token().secret().to_owned(), - refresh_token: token.refresh_token().map(|s| s.secret().to_owned()), - expires_in: token.expires_in(), - }) - } - - pub async fn get_user_info( - &self, - provider: &P, - access_token: &str, - ) -> Result { - let mut user_info_request = self - .http_client - .get(provider.user_info_url()) - .bearer_auth(access_token); - for (name, val) in provider.create_request_headers() { - user_info_request = user_info_request.header(name, val); - } - - let user_info_val = user_info_request - .send() - .await? - .error_for_status()? - .json() - .await?; - let user_info = provider.extract_user_info(user_info_val)?; - - Ok(user_info) - } -} diff --git a/server-new/crates/simple-oauth/src/provider.rs b/server-new/crates/simple-oauth/src/provider.rs deleted file mode 100644 index c764703..0000000 --- a/server-new/crates/simple-oauth/src/provider.rs +++ /dev/null @@ -1,15 +0,0 @@ -use std::fmt::Debug; - -use crate::types::UserInfo; - -/// Trait for all OAuth providers -pub trait SimpleOAuthProvider: Debug + Send + Sync { - fn authorize_url(&self) -> &str; - fn token_url(&self) -> &str; - fn user_info_url(&self) -> &str; - fn default_scopes(&self) -> &'static [&'static str]; - fn create_request_headers(&self) -> Vec<(String, String)> { - vec![] - } - fn extract_user_info(&self, val: serde_json::Value) -> Result; -} diff --git a/server-new/crates/simple-oauth/src/types.rs b/server-new/crates/simple-oauth/src/types.rs deleted file mode 100644 index 8b163ee..0000000 --- a/server-new/crates/simple-oauth/src/types.rs +++ /dev/null @@ -1,83 +0,0 @@ -use std::{borrow::Cow, fmt::Debug}; - -use serde::Deserialize; - -/// OAuth2 authorization redirect URL, along with the state and PKCE verifier -pub struct AuthorizeUrl { - pub url: oauth2::url::Url, - pub state: String, - pub pkce_verifier: String, -} -impl Debug for AuthorizeUrl { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("AuthorizeUrl") - .field("url", &"--redacted--") - .field("state", &self.state) - .field("pkce_verifier", &"--redacted--") - .finish() - } -} - -/// User info returned by the OAuth provider -#[derive(Debug)] -pub struct UserInfo { - /// The ID of the user at the OAuth provider - pub id: String, - /// The user's display name - pub name: Option, - /// The user's email. Likely will not be included unless you add the proper email scope for the provider. - /// - /// ⚠️ Do not rely on this for identifying the user. Use the `id` and the provider name. - pub email: Option, - /// Whether the user's email is verified. Not all providers return this in the user info. - pub email_verified: Option, - /// The URL of the user's picture/avatar - pub avatar_url: Option, -} - -/// Standard OAuth2 token response -pub struct StandardTokenResponse { - pub access_token: String, - pub refresh_token: Option, - pub expires_in: Option, -} -impl Debug for StandardTokenResponse { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("StandardTokenResponse") - .field("access_token", &"--redacted--") - .field("refresh_token", &"--redacted--") - .field("expires_in", &self.expires_in) - .finish() - } -} - -pub struct OAuthCredentials<'a> { - pub client_id: Cow<'a, str>, - pub client_secret: Cow<'a, str>, -} -impl<'a> OAuthCredentials<'a> { - pub fn new(client_id: impl Into>, client_secret: impl Into>) -> Self { - Self { - client_id: client_id.into(), - client_secret: client_secret.into(), - } - } -} -impl<'a> Debug for OAuthCredentials<'a> { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("OAuthCredentials") - .field("client_id", &self.client_id) - .field("client_secret", &"--redacted--") - .finish() - } -} - -/// OIDC discovery document -#[derive(Debug, Default, Deserialize)] -pub struct OidcDiscovery { - pub issuer: String, - pub authorization_endpoint: String, - pub token_endpoint: String, - pub userinfo_endpoint: String, - pub scopes_supported: Option>, -} diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index 30787bd..bc5ab45 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -2,9 +2,8 @@ use futures::future::BoxFuture; use serde::{Deserialize, Serialize}; use simple_oauth::{ SimpleOAuthClient, SimpleOAuthProvider, - types::{AuthorizeUrl, OAuthCredentials, StandardTokenResponse, UserInfo}, + types::{OAuthCredentials, StandardTokenResponse, UserInfo}, }; -use subtle::ConstantTimeEq; use tower_sessions::Session; use crate::{ @@ -45,7 +44,7 @@ impl OAuthProviderEnum { /// Trait for all OAuth providers pub trait OAuthProvider: Send + Sync { fn get_inner_provider(&self) -> Box; - fn get_credentials(&self) -> OAuthCredentials<'_>; + fn get_credentials(&self) -> OAuthCredentials; fn find_linked_user<'a>( &self, db: &'a mut DbService, @@ -85,22 +84,18 @@ impl<'a> OAuthService<'a> { callback_path: &str, session: &Session, ) -> AuthResult { - let provider = self.get_provider(provider)?; - let client = SimpleOAuthClient::with_http_client(self.http_client.clone()); - let AuthorizeUrl { - url, - state, - pkce_verifier, - } = client.authorize_url( - provider.get_inner_provider().as_ref(), - provider.get_credentials(), - &self.get_redirect_url(callback_path), - None, - )?; - session.insert(Self::SESS_STATE_FIELD, state).await?; - session.insert(Self::SESS_PKCE_FIELD, pkce_verifier).await?; - - Ok(url) + let (client, _) = self.get_client(provider)?; + let auth = client + .authorize_url() + .redirect_url(self.get_redirect_url(callback_path)) + .build()?; + + session.insert(Self::SESS_STATE_FIELD, auth.state).await?; + session + .insert(Self::SESS_PKCE_FIELD, auth.pkce_verifier) + .await?; + + Ok(auth.url) } pub async fn exchange_code( @@ -112,7 +107,7 @@ impl<'a> OAuthService<'a> { state: &str, ) -> AuthResult { // Get saved state and code verifier from session - let saved_state = session + let initial_state = session .remove::(Self::SESS_STATE_FIELD) .await? .ok_or(AuthError::Unauthorized("missing state in session"))?; @@ -121,22 +116,16 @@ impl<'a> OAuthService<'a> { .await? .ok_or(AuthError::Unauthorized("missing PKCE in session"))?; - // Verify state - if saved_state.as_bytes().ct_ne(state.as_bytes()).into() { - return Err(AuthError::Unauthorized("OAuth state mismatch")); - } - // Exchange code for token - let provider = self.get_provider(provider)?; - let client = SimpleOAuthClient::with_http_client(self.http_client.clone()); + let (client, _) = self.get_client(provider)?; let response = client - .exchange_code( - provider.get_inner_provider().as_ref(), - provider.get_credentials(), - &self.get_redirect_url(callback_path), - code, - Some(&pkce_verifier), - ) + .exchange_code() + .redirect_url(self.get_redirect_url(callback_path)) + .code(code) + .state(state) + .initial_state(&initial_state) + .pkce_verifier(pkce_verifier) + .build() .await?; Ok(response) @@ -148,13 +137,9 @@ impl<'a> OAuthService<'a> { token: &StandardTokenResponse, active_session: Option, ) -> AuthResult { - let provider = self.get_provider(provider)?; - let client = SimpleOAuthClient::with_http_client(self.http_client.clone()); - // Get user info from provider - let user_info = client - .get_user_info(provider.get_inner_provider().as_ref(), &token.access_token) - .await?; + let (client, provider) = self.get_client(provider)?; + let user_info = client.get_user_info(&token.access_token).await?; // Check for existing user, or create new user let mut db = DbService::from_pool(self.db).await?; @@ -195,7 +180,13 @@ impl<'a> OAuthService<'a> { format!("{}{}", &self.config.server.base_url, callback_path) } - fn get_provider(&self, provider: OAuthProviderEnum) -> AuthResult> { + fn get_client( + &self, + provider: OAuthProviderEnum, + ) -> AuthResult<( + SimpleOAuthClient>, + Box, + )> { let provider: Option> = match provider { OAuthProviderEnum::Github => match self.config.auth.github { Some(ref c) => Some(Box::new(github::GitHubOAuthProvider::new(c))), @@ -211,6 +202,13 @@ impl<'a> OAuthService<'a> { }, }; - provider.ok_or(AuthError::BadRequest("unsupported OAuth provider")) + let provider = provider.ok_or(AuthError::BadRequest("unsupported OAuth provider"))?; + let client = SimpleOAuthClient::builder() + .provider(provider.get_inner_provider()) + .credentials(provider.get_credentials()) + .http_client(self.http_client) + .build()?; + + Ok((client, provider)) } } diff --git a/server-new/src/services/auth/oauth/discord.rs b/server-new/src/services/auth/oauth/discord.rs index f0875cd..04f5dcd 100644 --- a/server-new/src/services/auth/oauth/discord.rs +++ b/server-new/src/services/auth/oauth/discord.rs @@ -27,10 +27,10 @@ impl DiscordOAuthProvider { impl OAuthProvider for DiscordOAuthProvider { fn get_inner_provider(&self) -> Box { - Box::new(simple_oauth::common::discord::Discord) + Box::new(simple_oauth::common::Discord) } - fn get_credentials(&self) -> OAuthCredentials<'_> { + fn get_credentials(&self) -> OAuthCredentials { OAuthCredentials::new( self.config.client_id.to_string(), &self.config.client_secret, @@ -62,7 +62,11 @@ impl OAuthProvider for DiscordOAuthProvider { fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { discord_id: Some(&user_data.id), - name: &user_data.name.as_deref().unwrap_or_default(), + name: &user_data + .name + .as_deref() + .or(user_data.username.as_deref()) + .unwrap_or_default(), avatar_url: user_data.avatar_url.as_deref(), ..Default::default() } diff --git a/server-new/src/services/auth/oauth/github.rs b/server-new/src/services/auth/oauth/github.rs index 060379a..c3041a0 100644 --- a/server-new/src/services/auth/oauth/github.rs +++ b/server-new/src/services/auth/oauth/github.rs @@ -30,10 +30,10 @@ impl GitHubOAuthProvider { impl OAuthProvider for GitHubOAuthProvider { fn get_inner_provider(&self) -> Box { - Box::new(simple_oauth::common::github::GitHub) + Box::new(simple_oauth::common::GitHub) } - fn get_credentials(&self) -> OAuthCredentials<'_> { + fn get_credentials(&self) -> OAuthCredentials { OAuthCredentials::new(&self.config.client_id, &self.config.client_secret) } @@ -62,7 +62,11 @@ impl OAuthProvider for GitHubOAuthProvider { fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { github_id: Some(&user_data.id), - name: &user_data.name.as_deref().unwrap_or_default(), + name: &user_data + .name + .as_deref() + .or(user_data.username.as_deref()) + .unwrap_or_default(), avatar_url: user_data.avatar_url.as_deref(), ..Default::default() } diff --git a/server-new/src/services/auth/oauth/google.rs b/server-new/src/services/auth/oauth/google.rs index a48901d..eb0b473 100644 --- a/server-new/src/services/auth/oauth/google.rs +++ b/server-new/src/services/auth/oauth/google.rs @@ -27,10 +27,10 @@ impl GoogleOAuthProvider { impl OAuthProvider for GoogleOAuthProvider { fn get_inner_provider(&self) -> Box { - Box::new(simple_oauth::common::google::Google) + Box::new(simple_oauth::common::Google) } - fn get_credentials(&self) -> OAuthCredentials<'_> { + fn get_credentials(&self) -> OAuthCredentials { OAuthCredentials::new(&self.config.client_id, &self.config.client_secret) } @@ -59,7 +59,11 @@ impl OAuthProvider for GoogleOAuthProvider { fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { google_id: Some(&user_data.id), - name: &user_data.name.as_deref().unwrap_or_default(), + name: &user_data + .name + .as_deref() + .or(user_data.username.as_deref()) + .unwrap_or_default(), avatar_url: user_data.avatar_url.as_deref(), ..Default::default() } From 2539d7e9432674f913b66b11dcb5cbf12b8a93cc Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sun, 28 Jun 2026 17:59:40 -0400 Subject: [PATCH 043/111] add OIDC provider --- server-new/src/config.rs | 3 +- server-new/src/db/repositories/user.rs | 18 ++-- server-new/src/lib.rs | 2 +- .../src/plugins/{session.rs => auth.rs} | 13 ++- server-new/src/plugins/mod.rs | 2 +- server-new/src/services/auth/error.rs | 2 +- server-new/src/services/auth/mod.rs | 15 ++- server-new/src/services/auth/oauth.rs | 99 ++++++++++++------- server-new/src/services/auth/oauth/discord.rs | 6 +- server-new/src/services/auth/oauth/github.rs | 6 +- server-new/src/services/auth/oauth/google.rs | 6 +- server-new/src/services/auth/oauth/oidc.rs | 91 +++++++++++++++++ server-new/src/state.rs | 14 ++- 13 files changed, 211 insertions(+), 66 deletions(-) rename server-new/src/plugins/{session.rs => auth.rs} (83%) create mode 100644 server-new/src/services/auth/oauth/oidc.rs diff --git a/server-new/src/config.rs b/server-new/src/config.rs index a0aa96d..76af488 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -6,7 +6,7 @@ use figment::providers::{Env, Format, Toml}; use serde::Deserialize; use crate::{ - services::auth::oauth::{DiscordOAuthConfig, GitHubOAuthConfig, GoogleOAuthConfig}, + services::auth::oauth::{DiscordOAuthConfig, GitHubOAuthConfig, GoogleOAuthConfig, OidcConfig}, state::AppState, }; @@ -43,6 +43,7 @@ pub struct AuthConfig { pub github: Option, pub discord: Option, pub google: Option, + pub oidc: Option, } #[derive(Debug, Clone, Deserialize)] diff --git a/server-new/src/db/repositories/user.rs b/server-new/src/db/repositories/user.rs index ad58f45..abb4908 100644 --- a/server-new/src/db/repositories/user.rs +++ b/server-new/src/db/repositories/user.rs @@ -70,16 +70,16 @@ impl<'a> UserRepository<'a> { Ok(user) } - // pub async fn find_by_oidc_id(&mut self, id: &str) -> Result, Error> { - // let user = users::table - // .filter(users::oidc_id.eq(id)) - // .select(ChatRsUser::as_select()) - // .first(self.db) - // .await - // .optional()?; + pub async fn find_by_oidc_id(&mut self, id: &str) -> Result, Error> { + let user = users::table + .filter(users::oidc_id.eq(id)) + .select(ChatRsUser::as_select()) + .first(self.db) + .await + .optional()?; - // Ok(user) - // } + Ok(user) + } // pub async fn find_by_sso_username(&mut self, username: &str) -> Result, Error> { // let user_id = users::table diff --git a/server-new/src/lib.rs b/server-new/src/lib.rs index e15dd49..bedeac0 100644 --- a/server-new/src/lib.rs +++ b/server-new/src/lib.rs @@ -21,7 +21,7 @@ pub async fn create_app() -> anyhow::Result> { .register(plugins::database::plugin()) // Initialize database .register(plugins::redis::plugin()) // Initialize Redis .register(api::plugin()) // Add API routes - .register(plugins::session::plugin()) // Setup sessions + .register(plugins::auth::plugin()) // Setup auth & sessions .register(plugins::logging::plugin()) // Request logging .register(plugins::security::plugin()) // Body limit, security headers, etc. .init() diff --git a/server-new/src/plugins/session.rs b/server-new/src/plugins/auth.rs similarity index 83% rename from server-new/src/plugins/session.rs rename to server-new/src/plugins/auth.rs index 5a49277..8d0f6cd 100644 --- a/server-new/src/plugins/session.rs +++ b/server-new/src/plugins/auth.rs @@ -6,20 +6,27 @@ use tower_sessions::{CachingSessionStore, Expiry, SessionManagerLayer, cookie}; use tower_sessions_redis_store::RedisStore; use crate::{ + config::AppConfig, db::DbPool, - services::auth::{session::AuthSessionService, session_store::SessionDbStore}, + services::auth::{ + oauth::OAuthService, session::AuthSessionService, session_store::SessionDbStore, + }, state::AppState, }; const REDIS_PREFIX: &str = "rs-chat:sess:"; const CLEANUP_INTERVAL: Duration = Duration::from_mins(15); -/// Add session handling to the server. Sessions are stored in Postgres and cached in Redis. +/// Add auth & session handling to the server. Sessions are stored in Postgres and cached in Redis. pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Session") - .on_init(async |state| { + .on_init(async |mut state| { + let config = state.get::().context("no config")?; let db_pool = state.get::().context("no db pool")?.clone(); + // Build configured OAuth providers + state.insert(OAuthService::build_provider_map(&config.auth)); + // Session cleanup task tokio::task::spawn(async move { let mut interval = tokio::time::interval(CLEANUP_INTERVAL); diff --git a/server-new/src/plugins/mod.rs b/server-new/src/plugins/mod.rs index dfe31c0..de05e46 100644 --- a/server-new/src/plugins/mod.rs +++ b/server-new/src/plugins/mod.rs @@ -1,5 +1,5 @@ +pub mod auth; pub mod database; pub mod logging; pub mod redis; pub mod security; -pub mod session; diff --git a/server-new/src/services/auth/error.rs b/server-new/src/services/auth/error.rs index b2e7854..d03b019 100644 --- a/server-new/src/services/auth/error.rs +++ b/server-new/src/services/auth/error.rs @@ -15,7 +15,7 @@ pub enum AuthError { UserNotFound, #[error("database error: {0}")] Database(#[from] diesel::result::Error), - #[error("database error: {0}")] + #[error("database pool error: {0}")] DatabasePool(#[from] DbPoolError), #[error("session error: {0}")] Session(#[from] tower_sessions::session::Error), diff --git a/server-new/src/services/auth/mod.rs b/server-new/src/services/auth/mod.rs index f93dd63..6ae5d2b 100644 --- a/server-new/src/services/auth/mod.rs +++ b/server-new/src/services/auth/mod.rs @@ -14,17 +14,24 @@ pub use error::{AuthError, AuthResult}; pub use types::*; pub struct AuthService<'a> { - config: &'a AppConfig, db: &'a DbPool, + config: &'a AppConfig, http_client: &'a reqwest::Client, + oauth_providers: &'a oauth::OAuthProviderMap, } impl<'a> AuthService<'a> { - pub fn new(config: &'a AppConfig, http_client: &'a reqwest::Client, db: &'a DbPool) -> Self { + pub fn new( + db: &'a DbPool, + config: &'a AppConfig, + http_client: &'a reqwest::Client, + oauth_providers: &'a oauth::OAuthProviderMap, + ) -> Self { Self { + db, config, http_client, - db, + oauth_providers, } } @@ -45,6 +52,6 @@ impl<'a> AuthService<'a> { /// Access OAuth functions pub fn oauth(self) -> oauth::OAuthService<'a> { - oauth::OAuthService::new(self.config, self.db, self.http_client) + oauth::OAuthService::new(self.config, self.db, self.http_client, self.oauth_providers) } } diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index bc5ab45..912955b 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -1,3 +1,5 @@ +use std::collections::HashMap; + use futures::future::BoxFuture; use serde::{Deserialize, Serialize}; use simple_oauth::{ @@ -18,10 +20,15 @@ use crate::{ mod discord; mod github; mod google; +mod oidc; pub use discord::DiscordOAuthConfig; pub use github::GitHubOAuthConfig; pub use google::GoogleOAuthConfig; +pub use oidc::OidcConfig; + +/// Map of configured OAuth providers stored in state +pub type OAuthProviderMap = HashMap>; /// Supported OAuth providers #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] @@ -30,6 +37,7 @@ pub enum OAuthProviderEnum { Github, Discord, Google, + Oidc, } impl OAuthProviderEnum { pub fn as_str(&self) -> &str { @@ -37,6 +45,7 @@ impl OAuthProviderEnum { OAuthProviderEnum::Github => "github", OAuthProviderEnum::Discord => "discord", OAuthProviderEnum::Google => "google", + OAuthProviderEnum::Oidc => "oidc", } } } @@ -60,6 +69,7 @@ pub struct OAuthService<'a> { config: &'a AppConfig, db: &'a DbPool, http_client: &'a reqwest::Client, + provider_map: &'a OAuthProviderMap, } impl<'a> OAuthService<'a> { @@ -70,11 +80,13 @@ impl<'a> OAuthService<'a> { config: &'a AppConfig, db: &'a DbPool, http_client: &'a reqwest::Client, + provider_map: &'a OAuthProviderMap, ) -> Self { Self { config, db, http_client, + provider_map, } } @@ -84,8 +96,9 @@ impl<'a> OAuthService<'a> { callback_path: &str, session: &Session, ) -> AuthResult { - let (client, _) = self.get_client(provider)?; - let auth = client + let oauth_provider = self.oauth_provider(provider)?; + let auth = self + .oauth_client(oauth_provider)? .authorize_url() .redirect_url(self.get_redirect_url(callback_path)) .build()?; @@ -117,8 +130,9 @@ impl<'a> OAuthService<'a> { .ok_or(AuthError::Unauthorized("missing PKCE in session"))?; // Exchange code for token - let (client, _) = self.get_client(provider)?; - let response = client + let oauth_provider = self.oauth_provider(provider)?; + let response = self + .oauth_client(oauth_provider)? .exchange_code() .redirect_url(self.get_redirect_url(callback_path)) .code(code) @@ -138,12 +152,15 @@ impl<'a> OAuthService<'a> { active_session: Option, ) -> AuthResult { // Get user info from provider - let (client, provider) = self.get_client(provider)?; - let user_info = client.get_user_info(&token.access_token).await?; + let oauth_provider = self.oauth_provider(provider)?; + let user_info = self + .oauth_client(oauth_provider)? + .get_user_info(&token.access_token) + .await?; // Check for existing user, or create new user let mut db = DbService::from_pool(self.db).await?; - let user = match provider.find_linked_user(&mut db, &user_info).await? { + let user = match oauth_provider.find_linked_user(&mut db, &user_info).await? { Some(existing_user) => { if active_session.is_some_and(|sess| sess.user_id != existing_user.id) { return Err(AuthError::Unauthorized("cannot switch users via OAuth")); @@ -153,16 +170,16 @@ impl<'a> OAuthService<'a> { } None => match active_session { None => { - let new_user = provider.create_new_user(&user_info); + let new_user = oauth_provider.create_new_user(&user_info); db.users().create(new_user).await? } Some(sess) => match db.users().find_by_id(&sess.user_id).await? { - Some(user) if provider.is_user_linked(&user) => { + Some(user) if oauth_provider.is_user_linked(&user) => { return Err(AuthError::BadRequest("user already linked to provider")); } Some(user) => { // Link logged-in user to new provider - let update_user = provider.create_update_user(&user_info); + let update_user = oauth_provider.create_update_user(&user_info); db.users().update(&user.id, update_user).await?; user } @@ -180,35 +197,47 @@ impl<'a> OAuthService<'a> { format!("{}{}", &self.config.server.base_url, callback_path) } - fn get_client( - &self, - provider: OAuthProviderEnum, - ) -> AuthResult<( - SimpleOAuthClient>, - Box, - )> { - let provider: Option> = match provider { - OAuthProviderEnum::Github => match self.config.auth.github { - Some(ref c) => Some(Box::new(github::GitHubOAuthProvider::new(c))), - None => None, - }, - OAuthProviderEnum::Discord => match self.config.auth.discord { - Some(ref c) => Some(Box::new(discord::DiscordOAuthProvider::new(c))), - None => None, - }, - OAuthProviderEnum::Google => match self.config.auth.google { - Some(ref c) => Some(Box::new(google::GoogleOAuthProvider::new(c))), - None => None, - }, - }; + fn oauth_provider(&self, provider: OAuthProviderEnum) -> AuthResult<&dyn OAuthProvider> { + let provider = self + .provider_map + .get(&provider) + .ok_or_else(|| AuthError::BadRequest("unsupported OAuth provider"))?; + Ok(provider.as_ref()) + } - let provider = provider.ok_or(AuthError::BadRequest("unsupported OAuth provider"))?; - let client = SimpleOAuthClient::builder() + fn oauth_client( + &self, + provider: &dyn OAuthProvider, + ) -> Result>, AuthError> { + Ok(simple_oauth::SimpleOAuthClient::builder() .provider(provider.get_inner_provider()) .credentials(provider.get_credentials()) .http_client(self.http_client) - .build()?; + .build()?) + } + + pub fn build_provider_map(config: &crate::config::AuthConfig) -> OAuthProviderMap { + use { + discord::DiscordProvider, github::GitHubProvider, google::GoogleProvider, + oidc::OidcProvider, + }; + let mut map: OAuthProviderMap = HashMap::new(); + if let Some(ref c) = config.github { + map.insert(OAuthProviderEnum::Github, Box::new(GitHubProvider::new(c))); + } + if let Some(ref c) = config.discord { + map.insert( + OAuthProviderEnum::Discord, + Box::new(DiscordProvider::new(c)), + ); + } + if let Some(ref c) = config.google { + map.insert(OAuthProviderEnum::Google, Box::new(GoogleProvider::new(c))); + } + if let Some(ref c) = config.oidc { + map.insert(OAuthProviderEnum::Oidc, Box::new(OidcProvider::new(c))); + } - Ok((client, provider)) + map } } diff --git a/server-new/src/services/auth/oauth/discord.rs b/server-new/src/services/auth/oauth/discord.rs index 04f5dcd..c803803 100644 --- a/server-new/src/services/auth/oauth/discord.rs +++ b/server-new/src/services/auth/oauth/discord.rs @@ -13,11 +13,11 @@ pub struct DiscordOAuthConfig { client_secret: String, } -pub struct DiscordOAuthProvider { +pub struct DiscordProvider { config: DiscordOAuthConfig, } -impl DiscordOAuthProvider { +impl DiscordProvider { pub fn new(config: &DiscordOAuthConfig) -> Self { Self { config: config.clone(), @@ -25,7 +25,7 @@ impl DiscordOAuthProvider { } } -impl OAuthProvider for DiscordOAuthProvider { +impl OAuthProvider for DiscordProvider { fn get_inner_provider(&self) -> Box { Box::new(simple_oauth::common::Discord) } diff --git a/server-new/src/services/auth/oauth/github.rs b/server-new/src/services/auth/oauth/github.rs index c3041a0..092bb7a 100644 --- a/server-new/src/services/auth/oauth/github.rs +++ b/server-new/src/services/auth/oauth/github.rs @@ -16,11 +16,11 @@ pub struct GitHubOAuthConfig { client_secret: String, } -pub struct GitHubOAuthProvider { +pub struct GitHubProvider { config: GitHubOAuthConfig, } -impl GitHubOAuthProvider { +impl GitHubProvider { pub fn new(config: &GitHubOAuthConfig) -> Self { Self { config: config.clone(), @@ -28,7 +28,7 @@ impl GitHubOAuthProvider { } } -impl OAuthProvider for GitHubOAuthProvider { +impl OAuthProvider for GitHubProvider { fn get_inner_provider(&self) -> Box { Box::new(simple_oauth::common::GitHub) } diff --git a/server-new/src/services/auth/oauth/google.rs b/server-new/src/services/auth/oauth/google.rs index eb0b473..16d7520 100644 --- a/server-new/src/services/auth/oauth/google.rs +++ b/server-new/src/services/auth/oauth/google.rs @@ -13,11 +13,11 @@ pub struct GoogleOAuthConfig { client_secret: String, } -pub struct GoogleOAuthProvider { +pub struct GoogleProvider { config: GoogleOAuthConfig, } -impl GoogleOAuthProvider { +impl GoogleProvider { pub fn new(config: &GoogleOAuthConfig) -> Self { Self { config: config.clone(), @@ -25,7 +25,7 @@ impl GoogleOAuthProvider { } } -impl OAuthProvider for GoogleOAuthProvider { +impl OAuthProvider for GoogleProvider { fn get_inner_provider(&self) -> Box { Box::new(simple_oauth::common::Google) } diff --git a/server-new/src/services/auth/oauth/oidc.rs b/server-new/src/services/auth/oauth/oidc.rs new file mode 100644 index 0000000..a668b66 --- /dev/null +++ b/server-new/src/services/auth/oauth/oidc.rs @@ -0,0 +1,91 @@ +use futures::future::BoxFuture; +use serde::Deserialize; +use simple_oauth::{ + SimpleOAuthProvider, + types::{OAuthCredentials, OidcDiscovery}, +}; + +use crate::{ + db::models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, + services::auth::AuthResult, +}; + +use super::OAuthProvider; + +#[derive(Clone, Debug, Deserialize)] +pub struct OidcConfig { + name: Option, + client_id: String, + client_secret: String, + auth_endpoint: String, + token_endpoint: String, + userinfo_endpoint: String, +} + +pub struct OidcProvider { + config: OidcConfig, +} + +impl OidcProvider { + pub fn new(config: &OidcConfig) -> Self { + Self { + config: config.clone(), + } + } +} + +impl OAuthProvider for OidcProvider { + fn get_inner_provider(&self) -> Box { + Box::new(simple_oauth::common::Oidc::from_config(OidcDiscovery { + authorization_endpoint: self.config.auth_endpoint.clone(), + token_endpoint: self.config.token_endpoint.clone(), + userinfo_endpoint: self.config.userinfo_endpoint.clone(), + ..Default::default() + })) + } + + fn get_credentials(&self) -> OAuthCredentials { + OAuthCredentials::new(&self.config.client_id, &self.config.client_secret) + } + + fn find_linked_user<'a>( + &self, + db: &'a mut crate::db::DbService, + user_info: &'a simple_oauth::types::UserInfo, + ) -> BoxFuture<'a, AuthResult>> { + Box::pin(async move { + let user = db.users().find_by_oidc_id(&user_info.id).await?; + Ok(user) + }) + } + + fn is_user_linked(&self, user: &ChatRsUser) -> bool { + user.oidc_id.is_some() + } + + fn create_update_user<'a>( + &self, + user_info: &'a simple_oauth::types::UserInfo, + ) -> crate::db::models::UpdateChatRsUser<'a> { + UpdateChatRsUser { + oidc_id: Some(&user_info.id), + ..Default::default() + } + } + + fn create_new_user<'a>( + &self, + user_info: &'a simple_oauth::types::UserInfo, + ) -> crate::db::models::NewChatRsUser<'a> { + NewChatRsUser { + google_id: Some(&user_info.id), + name: &user_info + .name + .as_deref() + .or(user_info.username.as_deref()) + .unwrap_or_default(), + avatar_url: user_info.avatar_url.as_deref(), + ..Default::default() + } + } +} diff --git a/server-new/src/state.rs b/server-new/src/state.rs index ede4370..88654da 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -4,7 +4,11 @@ use std::{ops::Deref, sync::Arc}; use axum_plugin::{AppState, TypeMap}; -use crate::{config::AppConfig, db::DbPool, services::auth::AuthService}; +use crate::{ + config::AppConfig, + db::DbPool, + services::auth::{AuthService, oauth::OAuthProviderMap}, +}; /// App state stored in the Axum router #[derive(Clone)] @@ -16,11 +20,17 @@ pub struct AppStateInner { pub http_client: reqwest::Client, pub db_pool: DbPool, pub redis: fred::prelude::Pool, + pub oauth_providers: OAuthProviderMap, } impl AppState { pub fn auth_service(&self) -> AuthService<'_> { - AuthService::new(&self.config, &self.http_client, &self.db_pool) + AuthService::new( + &self.db_pool, + &self.config, + &self.http_client, + &self.oauth_providers, + ) } } From 1e69b36703a8645b9e0618df2f0e56e210ea4580 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 29 Jun 2026 02:21:07 -0400 Subject: [PATCH 044/111] chat service - checkpoint --- server-new/Cargo.lock | 43 +++- server-new/Cargo.toml | 5 +- server-new/src/services/chat/error.rs | 8 + server-new/src/services/chat/mod.rs | 43 ++++ server-new/src/services/llm/error.rs | 25 +++ server-new/src/services/llm/interface.rs | 24 +++ server-new/src/services/llm/mod.rs | 4 + server-new/src/services/llm/providers/mod.rs | 4 + .../src/services/llm/providers/openai/mod.rs | 190 +++++++++++++++++ .../services/llm/providers/openai/request.rs | 193 ++++++++++++++++++ .../services/llm/providers/openai/response.rs | 169 +++++++++++++++ .../src/services/llm/providers/utils.rs | 54 +++++ server-new/src/services/llm/types.rs | 95 +++++++++ server-new/src/services/mod.rs | 2 + server-new/src/state.rs | 9 +- 15 files changed, 865 insertions(+), 3 deletions(-) create mode 100644 server-new/src/services/chat/error.rs create mode 100644 server-new/src/services/chat/mod.rs create mode 100644 server-new/src/services/llm/error.rs create mode 100644 server-new/src/services/llm/interface.rs create mode 100644 server-new/src/services/llm/mod.rs create mode 100644 server-new/src/services/llm/providers/mod.rs create mode 100644 server-new/src/services/llm/providers/openai/mod.rs create mode 100644 server-new/src/services/llm/providers/openai/request.rs create mode 100644 server-new/src/services/llm/providers/openai/response.rs create mode 100644 server-new/src/services/llm/providers/utils.rs create mode 100644 server-new/src/services/llm/types.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index e8389ff..843fcb6 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -70,6 +70,28 @@ dependencies = [ "rustversion", ] +[[package]] +name = "async-stream" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b5a71a6f37880a80d1d7f19efd781e4b5de42c88f0722cc13bcb6cc2cfe8476" +dependencies = [ + "async-stream-impl", + "futures-core", + "pin-project-lite", +] + +[[package]] +name = "async-stream-impl" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7c24de15d275a1ecfd47a380fb4d5ec9bfe0933f309ed5e705b775596a3574d" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "async-trait" version = "0.1.89" @@ -248,7 +270,7 @@ version = "3.9.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6dee98b0db6a962de883bf5d20362dee4d7ca0d12fe39a7c6c73c844e1cd7c1f" dependencies = [ - "darling 0.21.3", + "darling 0.23.0", "ident_case", "prettyplease", "proc-macro2", @@ -1915,6 +1937,7 @@ dependencies = [ "base64", "bytes", "futures-core", + "futures-util", "http", "http-body", "http-body-util", @@ -1934,12 +1957,14 @@ dependencies = [ "sync_wrapper", "tokio", "tokio-rustls", + "tokio-util", "tower", "tower-http 0.6.11", "tower-service", "url", "wasm-bindgen", "wasm-bindgen-futures", + "wasm-streams", "web-sys", ] @@ -1981,6 +2006,7 @@ name = "rs-chat-api" version = "0.1.0" dependencies = [ "anyhow", + "async-stream", "async-trait", "axum", "axum-helmet", @@ -2002,6 +2028,8 @@ dependencies = [ "simple-oauth", "thiserror 2.0.18", "tokio", + "tokio-stream", + "tokio-util", "tower", "tower-http 0.7.0", "tower-sessions", @@ -3170,6 +3198,19 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "wasm-streams" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d1ec4f6517c9e11ae630e200b2b65d193279042e28edd4a2cda233e46670bbb" +dependencies = [ + "futures-util", + "js-sys", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + [[package]] name = "web-sys" version = "0.3.102" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 8d46425..8a7fcde 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -7,6 +7,7 @@ publish = false [dependencies] anyhow = "1.0.102" +async-stream = "0.3.6" async-trait = "0.1.89" axum = { version = "0.8.9", features = ["json", "query"] } axum-helmet = "1.0.2" @@ -42,7 +43,7 @@ hex = "0.4.3" reqwest = { version = "0.13.4", default-features = false, - features = ["default-tls", "json"] + features = ["default-tls", "json", "stream"] } serde = { version = "1.0.228", features = ["derive"] } serde_json = "1.0.150" @@ -62,6 +63,8 @@ tokio = { default-features = false, features = ["macros", "net", "rt", "rt-multi-thread", "signal"] } +tokio-stream = { version = "0.1.18", default-features = false } +tokio-util = { version = "0.7.18", features = ["io"] } tower = { version = "0.5", default-features = false } tower-http = { version = "0.7.0", diff --git a/server-new/src/services/chat/error.rs b/server-new/src/services/chat/error.rs new file mode 100644 index 0000000..c3b814b --- /dev/null +++ b/server-new/src/services/chat/error.rs @@ -0,0 +1,8 @@ +use crate::services::llm::error::LlmRequestError; + +/// Chat service errors +#[derive(Debug, thiserror::Error)] +pub enum ChatError { + #[error(transparent)] + Request(#[from] LlmRequestError), +} diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs new file mode 100644 index 0000000..3a04855 --- /dev/null +++ b/server-new/src/services/chat/mod.rs @@ -0,0 +1,43 @@ +use futures::Stream; + +use crate::services::{ + chat::error::ChatError, + llm::{ + interface::{LlmProvider, LlmStreamChunkResult}, + providers::OpenAIProvider, + types::{LlmChatOptions, LlmChatRequest, LlmMessage, LlmUserMessage}, + }, +}; + +mod error; + +pub struct ChatService<'a> { + http_client: &'a reqwest::Client, + redis: &'a fred::clients::Client, +} + +impl<'a> ChatService<'a> { + pub fn new(http_client: &'a reqwest::Client, redis: &'a fred::clients::Client) -> Self { + Self { http_client, redis } + } + + pub async fn test_chat(&self) -> Result, ChatError> { + let provider: Box = + Box::new(OpenAIProvider::new(self.http_client, self.redis, "", None)); + let messages = vec![LlmMessage::User(LlmUserMessage { + text: "Hello!".into(), + ..Default::default() + })]; + let request = LlmChatRequest { + messages: messages, + options: LlmChatOptions { + model: "gpt-5-mini".into(), + ..Default::default() + }, + }; + + let response = provider.stream_chat(&request).await?; + + Ok(response) + } +} diff --git a/server-new/src/services/llm/error.rs b/server-new/src/services/llm/error.rs new file mode 100644 index 0000000..dbc7f94 --- /dev/null +++ b/server-new/src/services/llm/error.rs @@ -0,0 +1,25 @@ +/// Errors that can occur in an LLM provider request +#[derive(Debug, thiserror::Error)] +pub enum LlmRequestError { + #[error("Provider error: {0}")] + Provider(String), +} + +/// Errors that can occur during LLM streaming +#[derive(Debug, thiserror::Error)] +pub enum LlmStreamError { + #[error("Provider error: {0}")] + Provider(String), + #[error("Failed to parse event: {0}")] + Parsing(#[from] serde_json::Error), + #[error("Failed to decode line: {0}")] + Decoding(#[from] tokio_util::codec::LinesCodecError), + #[error("Stream was cancelled")] + StreamCancelled, + // #[error("Redis error: {0}")] + // Redis(#[from] fred::error::Error), + // #[error("Tinistream error: {0}")] + // Tinistream(#[from] crate::stream::TiniError), + // #[error("Websocket error: {0}")] + // Websocket(#[from] reqwest_websocket::Error), +} diff --git a/server-new/src/services/llm/interface.rs b/server-new/src/services/llm/interface.rs new file mode 100644 index 0000000..23d77e7 --- /dev/null +++ b/server-new/src/services/llm/interface.rs @@ -0,0 +1,24 @@ +use futures::{future::BoxFuture, stream::BoxStream}; + +use super::{ + error::{LlmRequestError, LlmStreamError}, + types::{LlmChatRequest, LlmUsage}, +}; + +/// Trait that all LLM providers must implement +pub trait LlmProvider { + fn stream_chat<'r>(&'r self, request: &'r LlmChatRequest) -> LlmStreamingResponse<'r>; +} + +pub type LlmStreamingResponse<'r> = BoxFuture<'r, Result>; +pub type LlmStream = BoxStream<'static, LlmStreamChunkResult>; +pub type LlmStreamChunkResult = Result; + +/// A streaming chunk of data from the LLM provider +pub enum LlmStreamChunk { + Text(String), + Usage(LlmUsage), + // ToolCalls(Vec), + // PendingToolCall(LlmPendingToolCall), + // Images(Vec), +} diff --git a/server-new/src/services/llm/mod.rs b/server-new/src/services/llm/mod.rs new file mode 100644 index 0000000..b667ff8 --- /dev/null +++ b/server-new/src/services/llm/mod.rs @@ -0,0 +1,4 @@ +pub mod error; +pub mod interface; +pub mod providers; +pub mod types; diff --git a/server-new/src/services/llm/providers/mod.rs b/server-new/src/services/llm/providers/mod.rs new file mode 100644 index 0000000..c219eef --- /dev/null +++ b/server-new/src/services/llm/providers/mod.rs @@ -0,0 +1,4 @@ +mod openai; +mod utils; + +pub use openai::OpenAIProvider; diff --git a/server-new/src/services/llm/providers/openai/mod.rs b/server-new/src/services/llm/providers/openai/mod.rs new file mode 100644 index 0000000..cbcdf8e --- /dev/null +++ b/server-new/src/services/llm/providers/openai/mod.rs @@ -0,0 +1,190 @@ +//! OpenAI (and OpenAI compatible) LLM provider + +use futures::StreamExt; + +use crate::services::llm::{ + error::LlmRequestError, + interface::{LlmProvider, LlmStreamingResponse}, + providers::utils, + types::LlmChatRequest, +}; + +mod request; +mod response; + +use {request::*, response::*}; + +const OPENAI_API_BASE_URL: &str = "https://api.openai.com/v1"; +const OPENROUTER_API_BASE_URL: &str = "https://openrouter.ai/api/v1"; + +/// OpenAI chat provider +#[derive(Debug, Clone)] +pub struct OpenAIProvider { + client: reqwest::Client, + redis: fred::clients::Client, + api_key: String, + base_url: String, +} + +impl OpenAIProvider { + pub fn new( + http_client: &reqwest::Client, + redis: &fred::clients::Client, + api_key: &str, + base_url: Option<&str>, + ) -> Self { + Self { + client: http_client.clone(), + redis: redis.clone(), + api_key: api_key.to_owned(), + base_url: base_url.unwrap_or(OPENAI_API_BASE_URL).to_owned(), + } + } +} + +impl LlmProvider for OpenAIProvider { + fn stream_chat<'r>(&'r self, req: &'r LlmChatRequest) -> LlmStreamingResponse<'r> { + let openai_messages = build_openai_messages(&req.messages); + // let openai_tools = tools.as_ref().map(|t| build_openai_tools(t)); + // + let request = OpenAIRequest { + model: &req.options.model, + messages: openai_messages, + // OpenAI official API deprecated `max_tokens` for `max_completion_tokens` + max_tokens: match req.options.max_tokens { + Some(max_tokens) if self.base_url != OPENAI_API_BASE_URL => Some(max_tokens), + _ => None, + }, + max_completion_tokens: match req.options.max_tokens { + Some(max_tokens) if self.base_url == OPENAI_API_BASE_URL => Some(max_tokens), + _ => None, + }, + temperature: req.options.temperature, + store: (self.base_url == OPENAI_API_BASE_URL).then_some(false), + stream: Some(true), + stream_options: Some(OpenAIStreamOptions { + include_usage: true, + }), + // tools: openai_tools, + // modalities: options.modalities.as_ref(), + ..Default::default() + }; + + Box::pin(async move { + let response = self + .client + .post(format!("{}/chat/completions", self.base_url)) + .header("authorization", format!("Bearer {}", self.api_key)) + .header("content-type", "application/json") + .json(&request) + .send() + .await + .map_err(|e| LlmRequestError::Provider(format!("OpenAI request failed: {}", e)))?; + + if !response.status().is_success() { + let status = response.status(); + let error_text = response.text().await.unwrap_or_default(); + return Err(LlmRequestError::Provider(format!( + "OpenAI API error {}: {}", + status, error_text + ))); + } + + let stream = async_stream::stream! { + let mut sse_event_stream = utils::get_sse_events(response); + let mut tool_calls: Vec = Vec::new(); + while let Some(event) = sse_event_stream.next().await { + match event { + Ok(event) => { + for chunk in parse_openai_event(event, &mut tool_calls) { + yield chunk; + } + } + Err(e) => yield Err(e), + } + } + // if !tool_calls.is_empty() { + // if let Some(llm_tools) = tools { + // let converted = tool_calls + // .into_iter() + // .filter_map(|tc| tc.convert(&llm_tools)) + // .collect(); + // yield Ok(LlmStreamChunk::ToolCalls(converted)); + // } + // } + }; + + Ok(stream.boxed()) + }) + } + + // async fn prompt( + // &self, + // message: &str, + // options: &LlmProviderOptions, + // ) -> Result { + // let request = OpenAIRequest { + // model: &options.model, + // messages: vec![OpenAIMessage { + // role: "user", + // content: Some(vec![OpenAIContent::Text { text: message }]), + // ..Default::default() + // }], + // max_tokens: options.max_tokens, + // temperature: options.temperature, + // store: (self.base_url == OPENAI_API_BASE_URL).then_some(false), + // ..Default::default() + // }; + + // let response = self + // .client + // .post(format!("{}/chat/completions", self.base_url)) + // .header("authorization", format!("Bearer {}", self.api_key)) + // .header("content-type", "application/json") + // .json(&request) + // .send() + // .await + // .map_err(|e| LlmError::ProviderError(format!("OpenAI request failed: {}", e)))?; + + // if !response.status().is_success() { + // let status = response.status(); + // let error_text = response.text().await.unwrap_or_default(); + // return Err(LlmError::ProviderError(format!( + // "OpenAI API error {}: {}", + // status, error_text + // ))); + // } + + // let mut openai_response: OpenAIResponse = response + // .json() + // .await + // .map_err(|e| LlmError::ProviderError(format!("Failed to parse response: {}", e)))?; + + // let text = openai_response + // .choices + // .get_mut(0) + // .and_then(|choice| choice.message.as_mut()) + // .and_then(|message| message.content.take()) + // .ok_or(LlmError::NoResponse)?; + + // if let Some(usage) = openai_response.usage { + // let usage: LlmUsage = usage.into(); + // println!("Prompt usage: {:?}", usage); + // } + + // Ok(text) + // } + + // async fn list_models(&self) -> Result, LlmError> { + // let models = models::ModelsDevService::new(&self.redis, &self.client) + // .list_models({ + // match self.base_url.as_str() { + // OPENROUTER_API_BASE_URL => models::ModelsDevServiceProvider::OpenRouter, + // _ => models::ModelsDevServiceProvider::OpenAI, + // } + // }) + // .await?; + + // Ok(models) + // } +} diff --git a/server-new/src/services/llm/providers/openai/request.rs b/server-new/src/services/llm/providers/openai/request.rs new file mode 100644 index 0000000..8ef0256 --- /dev/null +++ b/server-new/src/services/llm/providers/openai/request.rs @@ -0,0 +1,193 @@ +use serde::Serialize; + +use crate::services::llm::{ + providers::utils, + types::{LlmFileType, LlmMessage}, +}; + +pub fn build_openai_messages<'a>(messages: &'a [LlmMessage]) -> Vec> { + messages + .iter() + .map(|message| match message { + LlmMessage::User(user_message) => { + let mut content = Vec::new(); + if !user_message.text.is_empty() { + content.push(OpenAIContent::Text { + text: &user_message.text, + }); + } + if let Some(ref files) = user_message.files { + content.extend(files.iter().map(|file| match file.file_type { + LlmFileType::Text => OpenAIContent::Text { + text: &file.content, + }, + LlmFileType::Image => OpenAIContent::ImageUrl { + image_url: OpenAIImageUrl { + url: utils::create_data_uri(&file.content_type, &file.content), + }, + }, + LlmFileType::Pdf => OpenAIContent::File { + file: OpenAIFile { + file_data: utils::create_data_uri( + &file.content_type, + &file.content, + ), + filename: &file.name, + }, + }, + })); + } + OpenAIMessage { + role: "user", + content: Some(content), + ..Default::default() + } + } + LlmMessage::Assistant(assistant_message) => { + // let tool_calls = assistant_message.tool_calls.as_ref().map(|tc| { + // tc.iter() + // .map(|tc| OpenAIToolCall { + // id: &tc.id, + // tool_type: "function", + // function: OpenAIToolCallFunction { + // name: &tc.tool_name, + // arguments: serde_json::to_string(&tc.parameters) + // .unwrap_or_default(), + // }, + // }) + // .collect() + // }); + OpenAIMessage { + role: "assistant", + content: (!assistant_message.text.is_empty()).then(|| { + vec![OpenAIContent::Text { + text: &assistant_message.text, + }] + }), + // tool_calls, + ..Default::default() + } + } + LlmMessage::System(text) => OpenAIMessage { + role: "system", + content: Some(vec![OpenAIContent::Text { text }]), + ..Default::default() + }, + // LlmMessage::Tool(tool_result) => OpenAIMessage { + // role: "tool", + // content: Some(vec![OpenAIContent::Text { + // text: &tool_result.content, + // }]), + // tool_call_id: Some(&tool_result.tool_call_id), + // ..Default::default() + // }, + }) + .collect() +} + +// pub fn build_openai_tools<'a>(tools: &'a [LlmTool]) -> Vec> { +// tools +// .iter() +// .map(|tool| OpenAITool { +// tool_type: "function", +// function: OpenAIToolFunction { +// name: &tool.name, +// description: &tool.description, +// parameters: &tool.input_schema, +// strict: true, +// }, +// }) +// .collect() +// } + +/// OpenAI API request body +#[derive(Debug, Default, Serialize)] +pub struct OpenAIRequest<'a> { + pub model: &'a str, + pub messages: Vec>, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_completion_tokens: Option, + pub temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub store: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub stream: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub stream_options: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub tools: Option>>, + // #[serde(skip_serializing_if = "Option::is_none")] + // pub modalities: Option<&'a Vec>, +} + +/// OpenAI API request stream options +#[derive(Debug, Serialize)] +pub struct OpenAIStreamOptions { + pub include_usage: bool, +} + +/// OpenAI API request message +#[derive(Debug, Default, Serialize)] +pub struct OpenAIMessage<'a> { + pub role: &'a str, + #[serde(skip_serializing_if = "Option::is_none")] + pub content: Option>>, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_call_id: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_calls: Option>>, +} + +#[derive(Debug, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum OpenAIContent<'a> { + Text { text: &'a str }, + ImageUrl { image_url: OpenAIImageUrl }, + File { file: OpenAIFile<'a> }, +} + +#[derive(Debug, Serialize)] +pub struct OpenAIImageUrl { + url: String, +} + +#[derive(Debug, Serialize)] +pub struct OpenAIFile<'a> { + file_data: String, + filename: &'a str, +} + +/// OpenAI tool definition +#[derive(Debug, Serialize)] +pub struct OpenAITool<'a> { + #[serde(rename = "type")] + tool_type: &'a str, + function: OpenAIToolFunction<'a>, +} + +/// OpenAI tool function definition +#[derive(Debug, Serialize)] +pub struct OpenAIToolFunction<'a> { + name: &'a str, + description: &'a str, + strict: bool, + parameters: &'a serde_json::Value, +} + +/// OpenAI tool call in messages +#[derive(Debug, Serialize)] +pub struct OpenAIToolCall<'a> { + id: &'a str, + #[serde(rename = "type")] + tool_type: &'a str, + function: OpenAIToolCallFunction<'a>, +} + +/// OpenAI tool call function in messages +#[derive(Debug, Serialize)] +pub struct OpenAIToolCallFunction<'a> { + name: &'a str, + arguments: String, +} diff --git a/server-new/src/services/llm/providers/openai/response.rs b/server-new/src/services/llm/providers/openai/response.rs new file mode 100644 index 0000000..0bf881f --- /dev/null +++ b/server-new/src/services/llm/providers/openai/response.rs @@ -0,0 +1,169 @@ +use serde::Deserialize; + +use crate::services::llm::{error::LlmStreamError, interface::LlmStreamChunk, types::LlmUsage}; + +/// Parse chunks from an OpenAI SSE event +pub fn parse_openai_event( + mut event: OpenAIStreamResponse, + tool_calls: &mut Vec, +) -> Vec> { + let mut chunks = Vec::with_capacity(1); + if let Some(delta) = event.choices.pop().and_then(|c| c.delta) { + if let Some(text) = delta.content { + chunks.push(Ok(LlmStreamChunk::Text(text))); + } + // if let Some(tool_calls_delta) = delta.tool_calls { + // for tool_call_delta in tool_calls_delta { + // if let Some(tc) = tool_calls + // .iter_mut() + // .find(|tc| tc.index == tool_call_delta.index) + // { + // if let Some(function_arguments) = tool_call_delta.function.arguments { + // *tc.function.arguments.get_or_insert_default() += &function_arguments; + // } + // if let Some(ref tool_name) = tc.function.name { + // let chunk = LlmStreamChunk::PendingToolCall(LlmPendingToolCall { + // index: tool_call_delta.index, + // tool_name: tool_name.clone(), + // }); + // chunks.push(Ok(chunk)); + // } + // } else { + // if let Some(ref tool_name) = tool_call_delta.function.name { + // let chunk = LlmStreamChunk::PendingToolCall(LlmPendingToolCall { + // index: tool_call_delta.index, + // tool_name: tool_name.clone(), + // }); + // chunks.push(Ok(chunk)); + // } + // tool_calls.push(tool_call_delta); + // } + // } + // } + // if let Some(images) = delta.images { + // chunks.push(Ok(LlmStreamChunk::Images( + // images + // .into_iter() + // .map(|image| LlmImage { + // base64_url: image.image_url.url, + // }) + // .collect(), + // ))); + // } + } + if let Some(usage) = event.usage { + chunks.push(Ok(LlmStreamChunk::Usage(usage.into()))); + } + + chunks +} + +/// OpenAI API response +#[derive(Debug, Deserialize)] +pub struct OpenAIResponse { + pub choices: Vec, + pub usage: Option, +} + +/// OpenAI API streaming response +#[derive(Debug, Deserialize)] +pub struct OpenAIStreamResponse { + choices: Vec, + usage: Option, +} + +/// OpenAI API response choice +#[derive(Debug, Deserialize)] +pub struct OpenAIChoice { + pub message: Option, + pub delta: Option, + // finish_reason: Option, +} + +/// OpenAI API response message +#[derive(Debug, Deserialize)] +pub struct OpenAIResponseMessage { + // role: String, + pub content: Option, +} + +/// OpenAI API streaming delta +#[derive(Debug, Deserialize)] +pub struct OpenAIResponseDelta { + // role: Option, + content: Option, + // tool_calls: Option>, + // /// OpenRouter images + // #[serde(skip_serializing_if = "Option::is_none")] + // pub images: Option>, +} + +/// OpenAI streaming tool call +#[derive(Debug, Deserialize)] +pub struct OpenAIStreamToolCall { + id: Option, + index: usize, + function: OpenAIStreamToolCallFunction, +} + +// impl OpenAIStreamToolCall { +// /// Convert OpenAI tool call format to ChatRsToolCall, add tool ID +// pub fn convert(self, rs_chat_tools: &[LlmTool]) -> Option { +// let id = self.id?; +// let tool_name = self.function.name?; +// let parameters = serde_json::from_str(&self.function.arguments?).ok()?; +// rs_chat_tools +// .iter() +// .find(|tool| tool.name == tool_name) +// .map(|tool| ChatRsToolCall { +// id, +// tool_id: tool.tool_id, +// tool_name, +// tool_type: tool.tool_type, +// parameters, +// }) +// } +// } + +/// OpenAI streaming tool call function +#[derive(Debug, Deserialize)] +struct OpenAIStreamToolCallFunction { + name: Option, + arguments: Option, +} + +// /// OpenRouter image +// #[derive(Debug, Deserialize)] +// pub struct OpenRouterImage { +// // #[serde(rename = "type")] +// // pub image_type: String, +// pub image_url: OpenRouterImageData, +// } + +// /// OpenRouter image data +// #[derive(Debug, Deserialize)] +// pub struct OpenRouterImageData { +// /// Base64 data URL +// pub url: String, +// } + +/// OpenAI API response usage +#[derive(Debug, Deserialize)] +pub struct OpenAIUsage { + prompt_tokens: Option, + completion_tokens: Option, + /// OpenRouter cost + cost: Option, + /// LLM Gateway cost + cost_usd_total: Option, +} + +impl From for LlmUsage { + fn from(usage: OpenAIUsage) -> Self { + LlmUsage { + input_tokens: usage.prompt_tokens, + output_tokens: usage.completion_tokens, + cost: usage.cost.or(usage.cost_usd_total), + } + } +} diff --git a/server-new/src/services/llm/providers/utils.rs b/server-new/src/services/llm/providers/utils.rs new file mode 100644 index 0000000..62ed993 --- /dev/null +++ b/server-new/src/services/llm/providers/utils.rs @@ -0,0 +1,54 @@ +//! Utilities for working with LLM requests and responses + +use futures::TryStreamExt; +use serde::de::DeserializeOwned; +use tokio_stream::{Stream, StreamExt}; +use tokio_util::{ + codec::{FramedRead, LinesCodec}, + io::StreamReader, +}; + +use crate::services::llm::error::LlmStreamError; + +/// Create a data URI +pub fn create_data_uri(content_type: &str, b64_string: &str) -> String { + format!("data:{content_type};base64,{b64_string}") +} + +/// Get a stream of deserialized events from a provider SSE stream. +pub fn get_sse_events( + response: reqwest::Response, +) -> impl Stream> { + let stream_reader = StreamReader::new(response.bytes_stream().map_err(std::io::Error::other)); + let line_reader = FramedRead::new(stream_reader, LinesCodec::new()); + + line_reader.filter_map(|line_result| { + match line_result { + Ok(line) => { + if line.len() >= 6 && line.as_bytes().starts_with(b"data: ") { + let data = &line[6..]; // Skip "data: " prefix + if data.trim_start().is_empty() || data == "[DONE]" { + None // Skip empty lines and termination markers + } else { + Some(serde_json::from_str::(data).map_err(LlmStreamError::Parsing)) + } + } else { + None // Ignore non-data lines + } + } + Err(e) => Some(Err(LlmStreamError::Decoding(e))), + } + }) +} + +/// Get a stream of deserialized events from a provider JSON stream, not SSE (e.g. Ollama uses this format). +pub fn get_json_events( + response: reqwest::Response, +) -> impl Stream> { + let stream_reader = StreamReader::new(response.bytes_stream().map_err(std::io::Error::other)); + let line_reader = FramedRead::new(stream_reader, LinesCodec::new()); + line_reader.map(|line_result| match line_result { + Ok(line) => serde_json::from_str::(&line).map_err(LlmStreamError::Parsing), + Err(e) => Err(LlmStreamError::Decoding(e)), + }) +} diff --git a/server-new/src/services/llm/types.rs b/server-new/src/services/llm/types.rs new file mode 100644 index 0000000..9603a3f --- /dev/null +++ b/server-new/src/services/llm/types.rs @@ -0,0 +1,95 @@ +use serde::{Deserialize, Serialize}; + +/// Generic chat request for all LLM providers +pub struct LlmChatRequest { + pub messages: Vec, + // tools: Option>, + pub options: LlmChatOptions, +} + +/// Generic message type to send to LLM providers +pub enum LlmMessage { + User(LlmUserMessage), + Assistant(LlmAssistantMessage), + System(String), + // Tool(LlmToolResult), +} + +/// Generic chat options for all LLM providers +#[derive(Clone, Debug, Default, Serialize, Deserialize)] +pub struct LlmChatOptions { + pub model: String, + pub temperature: Option, + pub max_tokens: Option, + // /// Only supported for OpenRouter + // #[serde(skip_serializing_if = "Option::is_none")] + // pub modalities: Option>, +} + +#[derive(Default)] +pub struct LlmUserMessage { + pub text: String, + pub files: Option>, +} + +pub struct LlmFileInput { + pub name: String, + pub file_type: LlmFileType, + pub content_type: String, + pub content: String, +} + +#[derive(Debug, PartialEq, Eq, Hash)] +pub enum LlmFileType { + Text, + Image, + Pdf, +} + +pub struct LlmAssistantMessage { + pub text: String, + // pub tool_calls: Option>, +} + +/// Usage stats from the LLM provider +#[derive(Debug, Default, Serialize, Deserialize)] +pub struct LlmUsage { + pub input_tokens: Option, + pub output_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cost: Option, +} + +// pub struct LlmToolCall { +// pub id: String, +// pub tool_id: Uuid, +// pub name: String, +// pub tool_type: LlmToolType, +// pub arguments: serde_json::Value, +// } + +// /// Generic tool that can be passed to LLM providers +// #[derive(Debug)] +// pub struct LlmTool { +// pub name: String, +// pub description: String, +// pub input_schema: serde_json::Value, +// /// ID of the RsChat tool that this is derived from +// pub tool_id: Uuid, +// /// The type of tool this is derived from (internal, external API, etc.) +// pub tool_type: LlmToolType, +// } + +// #[derive(Default, Debug, Clone, Copy, Serialize, Deserialize)] +// #[serde(rename_all = "snake_case")] +// pub enum LlmToolType { +// #[default] +// System, +// ExternalApi, +// } + +// pub struct LlmToolResult { +// pub tool_call_id: String, +// pub tool_name: String, +// pub content: String, +// } diff --git a/server-new/src/services/mod.rs b/server-new/src/services/mod.rs index 0e4a05d..0833d30 100644 --- a/server-new/src/services/mod.rs +++ b/server-new/src/services/mod.rs @@ -1 +1,3 @@ pub mod auth; +pub mod chat; +pub mod llm; diff --git a/server-new/src/state.rs b/server-new/src/state.rs index 88654da..bd34ce4 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -7,7 +7,10 @@ use axum_plugin::{AppState, TypeMap}; use crate::{ config::AppConfig, db::DbPool, - services::auth::{AuthService, oauth::OAuthProviderMap}, + services::{ + auth::{AuthService, oauth::OAuthProviderMap}, + chat::ChatService, + }, }; /// App state stored in the Axum router @@ -32,6 +35,10 @@ impl AppState { &self.oauth_providers, ) } + + pub fn chat_service(&self) -> ChatService<'_> { + ChatService::new(&self.http_client, self.redis.next()) + } } impl Deref for AppState { From 0c3740ee90ad5f2343142988129b6ba53d2f3da0 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 30 Jun 2026 01:51:39 -0400 Subject: [PATCH 045/111] chat service - checkpoint 2 --- server-new/Cargo.lock | 129 ++++++++- server-new/Cargo.toml | 6 + server-new/src/db/mod.rs | 37 ++- server-new/src/db/models.rs | 4 +- server-new/src/db/models/chat.rs | 140 ++++++++++ server-new/src/db/queries.rs | 75 ++++++ server-new/src/db/repositories.rs | 2 + server-new/src/db/repositories/chat.rs | 193 ++++++++++++++ server-new/src/lib.rs | 1 + server-new/src/{services => }/llm/error.rs | 14 +- .../src/{services => }/llm/interface.rs | 13 +- server-new/src/{services => }/llm/mod.rs | 2 + server-new/src/llm/providers/mod.rs | 4 + .../llm/providers/openai/mod.rs | 169 ++++++++++-- .../llm/providers/openai/request.rs | 2 +- .../llm/providers/openai/response.rs | 6 +- .../src/{services => }/llm/providers/utils.rs | 14 +- server-new/src/{services => }/llm/types.rs | 6 +- server-new/src/services/chat/error.rs | 12 +- server-new/src/services/chat/messages.rs | 76 ++++++ server-new/src/services/chat/mod.rs | 144 ++++++++-- server-new/src/services/llm/providers/mod.rs | 4 - server-new/src/services/mod.rs | 2 +- server-new/src/services/stream/error.rs | 8 + server-new/src/services/stream/mod.rs | 82 ++++++ server-new/src/services/stream/tests/mod.rs | 252 ++++++++++++++++++ server-new/src/services/stream/tests/utils.rs | 15 ++ server-new/src/services/stream/tinistream.rs | 161 +++++++++++ server-new/src/services/stream/writer.rs | 250 +++++++++++++++++ server-new/src/state.rs | 4 +- 30 files changed, 1724 insertions(+), 103 deletions(-) create mode 100644 server-new/src/db/models/chat.rs create mode 100644 server-new/src/db/queries.rs create mode 100644 server-new/src/db/repositories/chat.rs rename server-new/src/{services => }/llm/error.rs (60%) rename server-new/src/{services => }/llm/interface.rs (61%) rename server-new/src/{services => }/llm/mod.rs (58%) create mode 100644 server-new/src/llm/providers/mod.rs rename server-new/src/{services => }/llm/providers/openai/mod.rs (58%) rename server-new/src/{services => }/llm/providers/openai/request.rs (99%) rename server-new/src/{services => }/llm/providers/openai/response.rs (96%) rename server-new/src/{services => }/llm/providers/utils.rs (84%) rename server-new/src/{services => }/llm/types.rs (95%) create mode 100644 server-new/src/services/chat/messages.rs delete mode 100644 server-new/src/services/llm/providers/mod.rs create mode 100644 server-new/src/services/stream/error.rs create mode 100644 server-new/src/services/stream/mod.rs create mode 100644 server-new/src/services/stream/tests/mod.rs create mode 100644 server-new/src/services/stream/tests/utils.rs create mode 100644 server-new/src/services/stream/tinistream.rs create mode 100644 server-new/src/services/stream/writer.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index 843fcb6..0a5d9ab 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -103,6 +103,22 @@ dependencies = [ "syn", ] +[[package]] +name = "async-tungstenite" +version = "0.32.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8acc405d38be14342132609f06f02acaf825ddccfe76c4824a69281e0458ebd4" +dependencies = [ + "atomic-waker", + "futures-core", + "futures-io", + "futures-task", + "futures-util", + "log", + "pin-project-lite", + "tungstenite", +] + [[package]] name = "atomic" version = "0.6.1" @@ -539,6 +555,12 @@ dependencies = [ "syn", ] +[[package]] +name = "data-encoding" +version = "2.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4ae5f15dda3c708c0ade84bfee31ccab44a3da4f88015ed22f63732abe300c8" + [[package]] name = "deadpool" version = "0.13.0" @@ -597,6 +619,18 @@ dependencies = [ "tokio-postgres", ] +[[package]] +name = "diesel-derive-enum" +version = "3.0.0-beta.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "50a8c082045d01debc8589f8a0db9f2855a37c99c9b031325c856b5b98e1625f" +dependencies = [ + "heck 0.4.1", + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "diesel-jsonb-derive" version = "0.1.0" @@ -681,7 +715,7 @@ checksum = "dd122633e4bef06db27737f21d3738fb89c8f6d5360d6d9d7635dda142a7757e" dependencies = [ "darling 0.21.3", "either", - "heck", + "heck 0.5.0", "proc-macro2", "quote", "syn", @@ -960,6 +994,12 @@ version = "0.17.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" +[[package]] +name = "heck" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "95505c38b4572b2d910cecb0281560f54b440a19336cbbcb27bf6ce6adc6f5a8" + [[package]] name = "heck" version = "0.5.0" @@ -1752,6 +1792,21 @@ dependencies = [ "yansi", ] +[[package]] +name = "progenitor-client" +version = "0.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ffab7b358944dba033a7b324e7558e66e6bcb1fb4705cf57f26fd5092bcae630" +dependencies = [ + "bytes", + "futures-core", + "percent-encoding", + "reqwest", + "serde", + "serde_json", + "serde_urlencoded", +] + [[package]] name = "quinn" version = "0.11.11" @@ -1954,6 +2009,7 @@ dependencies = [ "rustls-platform-verifier", "serde", "serde_json", + "serde_urlencoded", "sync_wrapper", "tokio", "tokio-rustls", @@ -1968,6 +2024,26 @@ dependencies = [ "web-sys", ] +[[package]] +name = "reqwest-websocket" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7705b649c3b66b85c4e9c304a6898b1ae3eecb880c474720ebf925e4a932ae02" +dependencies = [ + "async-tungstenite", + "bytes", + "futures-util", + "reqwest", + "serde", + "serde_json", + "thiserror 2.0.18", + "tokio", + "tokio-util", + "tracing", + "tungstenite", + "web-sys", +] + [[package]] name = "ring" version = "0.17.14" @@ -2014,6 +2090,7 @@ dependencies = [ "chrono", "diesel", "diesel-async", + "diesel-derive-enum", "diesel-jsonb-derive", "diesel_migrations", "dotenvy", @@ -2022,11 +2099,13 @@ dependencies = [ "futures", "hex", "reqwest", + "reqwest-websocket", "serde", "serde_json", "serde_with", "simple-oauth", "thiserror 2.0.18", + "tinistream-client", "tokio", "tokio-stream", "tokio-util", @@ -2301,6 +2380,17 @@ dependencies = [ "syn", ] +[[package]] +name = "sha1" +version = "0.10.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" +dependencies = [ + "cfg-if", + "cpufeatures", + "digest", +] + [[package]] name = "sha2" version = "0.10.9" @@ -2551,6 +2641,19 @@ dependencies = [ "time-core", ] +[[package]] +name = "tinistream-client" +version = "0.1.10" +source = "git+https://github.com/fa-sharp/tinistream?rev=f25144c#f25144c1bdbee827d6033606b94a8aa1ae6eb5a7" +dependencies = [ + "bytes", + "futures-core", + "progenitor-client", + "reqwest", + "serde", + "serde_urlencoded", +] + [[package]] name = "tinystr" version = "0.8.3" @@ -2658,6 +2761,7 @@ checksum = "9ae9cec805b01e8fc3fd2fe289f89149a9b66dd16786abd8b19cfa7b48cb0098" dependencies = [ "bytes", "futures-core", + "futures-io", "futures-sink", "pin-project-lite", "tokio", @@ -2975,6 +3079,23 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" +[[package]] +name = "tungstenite" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8628dcc84e5a09eb3d8423d6cb682965dea9133204e8fb3efee74c2a0c259442" +dependencies = [ + "bytes", + "data-encoding", + "http", + "httparse", + "log", + "rand 0.9.4", + "sha1", + "thiserror 2.0.18", + "utf-8", +] + [[package]] name = "type-map" version = "0.5.1" @@ -3061,6 +3182,12 @@ version = "2.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "daf8dba3b7eb870caf1ddeed7bc9d2a049f3cfdfae7cb521b087cc33ae4c49da" +[[package]] +name = "utf-8" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9" + [[package]] name = "utf8_iter" version = "1.0.4" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 8a7fcde..3a540f1 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -29,6 +29,7 @@ diesel-async = { version = "0.9.2", features = ["deadpool", "migrations", "postgres"] } +diesel-derive-enum = { version = "3.0.0-beta.1", features = ["postgres"] } diesel-jsonb-derive = { path = "crates/diesel-jsonb-derive" } diesel_migrations = { version = "2.3.2", features = ["postgres"] } dotenvy = "0.15.7" @@ -45,6 +46,7 @@ reqwest = { default-features = false, features = ["default-tls", "json", "stream"] } +reqwest-websocket = { version = "0.6.0", features = ["json"] } serde = { version = "1.0.228", features = ["derive"] } serde_json = "1.0.150" serde_with = { @@ -58,6 +60,10 @@ simple-oauth = { features = ["default-tls"] } thiserror = "2.0.18" +tinistream-client = { + git = "https://github.com/fa-sharp/tinistream", + rev = "f25144c" +} tokio = { version = "1.52.3", default-features = false, diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs index 400bfe8..d597ace 100644 --- a/server-new/src/db/mod.rs +++ b/server-new/src/db/mod.rs @@ -1,14 +1,34 @@ +use std::ops::{Deref, DerefMut}; + +use diesel_async::{ + AsyncPgConnection, + pooled_connection::deadpool::{Object, Pool, PoolError}, +}; + pub mod models; +mod queries; mod repositories; mod schema; /// Type of the database pool -pub type DbPool = diesel_async::pooled_connection::deadpool::Pool; +pub type DbPool = Pool; /// Error when attempting to retrieve a connection from the pool -pub type DbPoolError = diesel_async::pooled_connection::deadpool::PoolError; -/// Type of the database connection retrieved from the pool -pub type DbConnection = - diesel_async::pooled_connection::deadpool::Object; +pub type DbPoolError = PoolError; + +/// The database connection retrieved from the pool. +/// For pipelining multiple queries, can use a shared reference with `&mut &**conn`. +pub struct DbConnection(Object); +impl Deref for DbConnection { + type Target = AsyncPgConnection; + fn deref(&self) -> &Self::Target { + self.0.as_ref() + } +} +impl DerefMut for DbConnection { + fn deref_mut(&mut self) -> &mut Self::Target { + &mut self.0 + } +} /// Date/time format used in all database tables pub type UtcDateTime = chrono::DateTime; @@ -23,17 +43,18 @@ impl DbService { pub fn new(cxn: DbConnection) -> Self { Self { cxn } } - pub async fn from_pool(pool: &DbPool) -> Result { let cxn = pool.get().await?; - Ok(Self::new(cxn)) + Ok(Self::new(DbConnection(cxn))) } pub fn users(&mut self) -> repositories::UserRepository<'_> { repositories::UserRepository::new(&mut self.cxn) } - pub fn sessions(&mut self) -> repositories::SessionRepository<'_> { repositories::SessionRepository::new(&mut self.cxn) } + pub fn chats(&mut self) -> repositories::ChatRepository<'_> { + repositories::ChatRepository::new(&mut self.cxn) + } } diff --git a/server-new/src/db/models.rs b/server-new/src/db/models.rs index c225868..80c3d61 100644 --- a/server-new/src/db/models.rs +++ b/server-new/src/db/models.rs @@ -1,7 +1,7 @@ use crate::db::schema; // mod api_key; -// mod chat; +mod chat; // mod file; // mod provider; // mod secret; @@ -10,7 +10,7 @@ mod session; mod user; // pub use api_key::*; -// pub use chat::*; +pub use chat::*; // pub use file::*; // pub use provider::*; // pub use secret::*; diff --git a/server-new/src/db/models/chat.rs b/server-new/src/db/models/chat.rs new file mode 100644 index 0000000..6301c17 --- /dev/null +++ b/server-new/src/db/models/chat.rs @@ -0,0 +1,140 @@ +use chrono::{DateTime, Utc}; +use diesel::{deserialize::FromSqlRow, expression::AsExpression, prelude::*}; +use diesel_jsonb_derive::AsJsonb; +use serde::{Deserialize, Serialize}; +use uuid::Uuid; + +use crate::{ + db::models::ChatRsUser, + llm::types::{LlmChatOptions, LlmUsage}, +}; + +#[derive(Identifiable, Associations, Queryable, Selectable, Serialize)] +#[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] +#[diesel(table_name = super::schema::chat_sessions)] +pub struct ChatRsSession { + pub id: Uuid, + #[serde(skip)] + pub user_id: Uuid, + pub title: String, + pub meta: ChatRsSessionMeta, + pub created_at: DateTime, + pub updated_at: DateTime, +} + +#[derive(Debug, Default, Serialize, Deserialize, FromSqlRow, AsExpression, AsJsonb)] +#[diesel(sql_type = diesel::sql_types::Jsonb)] +pub struct ChatRsSessionMeta { + // /// User configuration of tools for this session + // #[serde(skip_serializing_if = "Option::is_none")] + // pub tool_config: Option, +} +impl ChatRsSessionMeta { + pub fn new() -> Self { + Self {} + } +} + +#[derive(Insertable)] +#[diesel(table_name = super::schema::chat_sessions)] +pub struct NewChatRsSession<'r> { + pub user_id: &'r Uuid, + pub title: &'r str, +} + +#[derive(AsChangeset, Default)] +#[diesel(table_name = super::schema::chat_sessions)] +pub struct UpdateChatRsSession<'r> { + pub title: Option<&'r str>, + pub meta: Option, +} + +#[derive(diesel_derive_enum::DbEnum)] +#[db_enum(existing_type_path = "crate::db::schema::sql_types::ChatMessageRole")] +#[derive(Debug, PartialEq, Eq, Serialize)] +pub enum ChatRsMessageRole { + User, + Assistant, + System, + Tool, +} + +#[derive(Identifiable, Queryable, Selectable, Associations, Serialize)] +#[diesel(belongs_to(ChatRsSession, foreign_key = session_id))] +#[diesel(table_name = super::schema::chat_messages)] +pub struct ChatRsMessage { + pub id: Uuid, + pub session_id: Uuid, + pub role: ChatRsMessageRole, + pub content: String, + pub meta: ChatRsMessageMeta, + pub created_at: DateTime, +} + +#[derive(Debug, Default, Serialize, Deserialize, AsExpression, FromSqlRow, AsJsonb)] +#[diesel(sql_type = diesel::sql_types::Jsonb)] +pub struct ChatRsMessageMeta { + /// User messages: metadata associated with the user message + #[serde(skip_serializing_if = "Option::is_none")] + pub user: Option, + /// Assistant messages: metadata associated with the assistant message + #[serde(skip_serializing_if = "Option::is_none")] + pub assistant: Option, + // /// Tool messages: metadata of the executed tool call + // #[serde(skip_serializing_if = "Option::is_none")] + // pub tool_call: Option, +} +impl ChatRsMessageMeta { + pub fn new_assistant(assistant_meta: AssistantMeta) -> Self { + Self { + assistant: Some(assistant_meta), + ..Default::default() + } + } + pub fn new_user(user_meta: UserMeta) -> Self { + Self { + user: Some(user_meta), + ..Default::default() + } + } +} + +#[derive(Debug, Default, Serialize, Deserialize)] +pub struct UserMeta { + /// The IDs of the files attached to this message + #[serde(skip_serializing_if = "Option::is_none")] + pub files: Option>, +} + +#[derive(Debug, Default, Serialize, Deserialize)] +pub struct AssistantMeta { + /// The ID of the LLM provider used to generate this message + pub provider_id: i32, + /// Options passed to the LLM provider + #[serde(skip_serializing_if = "Option::is_none")] + pub provider_options: Option, + /// The tool calls requested by the assistant + // #[serde(skip_serializing_if = "Option::is_none")] + // pub tool_calls: Option>, + /// IDs of generated files + #[serde(skip_serializing_if = "Option::is_none")] + pub files: Option>, + /// Provider usage information + #[serde(skip_serializing_if = "Option::is_none")] + pub usage: Option, + /// Errors encountered during message generation + #[serde(skip_serializing_if = "Option::is_none")] + pub errors: Option>, + /// Whether this is a partial and/or interrupted message + #[serde(skip_serializing_if = "Option::is_none")] + pub partial: Option, +} + +#[derive(Insertable)] +#[diesel(table_name = super::schema::chat_messages)] +pub struct NewChatRsMessage<'r> { + pub session_id: &'r Uuid, + pub role: ChatRsMessageRole, + pub content: &'r str, + pub meta: ChatRsMessageMeta, +} diff --git a/server-new/src/db/queries.rs b/server-new/src/db/queries.rs new file mode 100644 index 0000000..3d8ea35 --- /dev/null +++ b/server-new/src/db/queries.rs @@ -0,0 +1,75 @@ +use diesel::{prelude::QueryableByName, sql_query}; +use diesel_async::RunQueryDsl; +use serde::Serialize; +use uuid::Uuid; + +use crate::db::DbConnection; + +/// Session matches for a full-text search query of chat titles and messages +#[derive(Debug, Clone, QueryableByName, Serialize)] +pub struct FullTextSearchResult { + #[diesel(sql_type = diesel::sql_types::Uuid)] + pub session_id: Uuid, + #[diesel(sql_type = diesel::sql_types::Double)] + pub session_rank: f64, + #[diesel(sql_type = diesel::sql_types::Timestamptz)] + pub session_updated_at: chrono::DateTime, + #[diesel(sql_type = diesel::sql_types::BigInt)] + pub message_matches: i64, + #[diesel(sql_type = diesel::sql_types::Text)] + pub title_highlight: String, + #[diesel(sql_type = diesel::sql_types::Text)] + pub message_highlights: String, +} + +/// Performs a full-text search of user's chat titles and messages +pub async fn full_text_query( + conn: &mut DbConnection, + user_id: &Uuid, + query: &str, + limit: i32, +) -> Result, diesel::result::Error> { + let results: Vec = sql_query( + r#" + WITH search_query AS ( + SELECT plainto_tsquery('english', $1) AS query + ), + message_stats AS ( + SELECT + cm.session_id, + cs.title, + cs.updated_at, + cm.content, + ts_rank(cm.search_vector, sq.query) AS rank, + COUNT(*) OVER (PARTITION BY cm.session_id) AS message_matches, + ROW_NUMBER() OVER ( + PARTITION BY cm.session_id + ORDER BY ts_rank(cm.search_vector, sq.query) DESC + ) AS rank_in_session + FROM chat_messages cm + JOIN chat_sessions cs ON cm.session_id = cs.id + CROSS JOIN search_query sq + WHERE cm.search_vector @@ sq.query + AND cs.user_id = $2 + ) + SELECT + session_id, + rank * (1 + LOG(message_matches) * 0.1) AS session_rank, + updated_at AS session_updated_at, + message_matches, + ts_headline('english', title, sq.query, 'StartSel=§§§HIGHLIGHT_START§§§, StopSel=§§§HIGHLIGHT_END§§§, HighlightAll=true') AS title_highlight, + ts_headline('english', content, sq.query, 'StartSel=§§§HIGHLIGHT_START§§§, StopSel=§§§HIGHLIGHT_END§§§, MinWords=8, MaxWords=12, MaxFragments=3') AS message_highlights + FROM message_stats ms + CROSS JOIN search_query sq + WHERE rank_in_session = 1 -- Only best message per session + ORDER BY session_rank DESC + LIMIT $3; + "#, + ) + .bind::(query) + .bind::(user_id) + .bind::(limit) + .load(conn).await?; + + Ok(results) +} diff --git a/server-new/src/db/repositories.rs b/server-new/src/db/repositories.rs index 7e3f251..959d916 100644 --- a/server-new/src/db/repositories.rs +++ b/server-new/src/db/repositories.rs @@ -1,5 +1,7 @@ +mod chat; mod session; mod user; +pub use chat::ChatRepository; pub use session::SessionRepository; pub use user::UserRepository; diff --git a/server-new/src/db/repositories/chat.rs b/server-new/src/db/repositories/chat.rs new file mode 100644 index 0000000..9de8566 --- /dev/null +++ b/server-new/src/db/repositories/chat.rs @@ -0,0 +1,193 @@ +use std::ops::{Deref, DerefMut}; + +use diesel::prelude::*; +use diesel_async::RunQueryDsl; +use uuid::Uuid; + +use crate::db::{ + DbConnection, + models::{ + ChatRsMessage, ChatRsSession, NewChatRsMessage, NewChatRsSession, UpdateChatRsSession, + }, + queries::{FullTextSearchResult, full_text_query}, + schema::{chat_messages, chat_sessions}, +}; + +pub struct ChatRepository<'a> { + pub db: &'a mut DbConnection, +} + +impl<'a> ChatRepository<'a> { + pub fn new(db: &'a mut DbConnection) -> Self { + ChatRepository { db } + } + + pub async fn create_session( + &mut self, + session: NewChatRsSession<'_>, + ) -> Result { + let id: Uuid = diesel::insert_into(chat_sessions::table) + .values(session) + .returning(chat_sessions::id) + .get_result(self.db) + .await?; + Ok(id.to_string()) + } + + pub async fn save_message( + &mut self, + message: NewChatRsMessage<'_>, + ) -> Result { + let message = diesel::insert_into(chat_messages::table) + .values(message) + .returning(ChatRsMessage::as_select()) + .get_result(self.db) + .await?; + Ok(message) + } + + pub async fn save_messages( + &mut self, + messages: &[NewChatRsMessage<'_>], + ) -> Result, diesel::result::Error> { + let messages = diesel::insert_into(chat_messages::table) + .values(messages) + .returning(ChatRsMessage::as_select()) + .get_results(self.db) + .await?; + Ok(messages) + } + + pub async fn find_message( + &mut self, + user_id: &Uuid, + message_id: &Uuid, + ) -> Result { + chat_messages::table + .inner_join(chat_sessions::table.on(chat_sessions::id.eq(chat_messages::session_id))) + .select(ChatRsMessage::as_select()) + .filter(chat_sessions::user_id.eq(user_id)) + .filter(chat_messages::id.eq(message_id)) + .get_result(self.db) + .await + } + + pub async fn delete_message( + &mut self, + session_id: &Uuid, + message_id: &Uuid, + ) -> Result { + let id: Uuid = diesel::delete(chat_messages::table) + .filter(chat_messages::session_id.eq(session_id)) + .filter(chat_messages::id.eq(message_id)) + .returning(chat_messages::id) + .get_result(self.db) + .await?; + Ok(id.to_string()) + } + + pub async fn get_all_sessions( + &mut self, + user_id: &Uuid, + ) -> Result, diesel::result::Error> { + let sessions = chat_sessions::table + .filter(chat_sessions::user_id.eq(user_id)) + .select(ChatRsSession::as_select()) + .order_by(chat_sessions::updated_at.desc()) + .limit(100) + .load(self.db) + .await?; + + Ok(sessions) + } + + pub async fn get_session( + &mut self, + user_id: &Uuid, + session_id: &Uuid, + ) -> Result { + let session = chat_sessions::table + .filter(chat_sessions::user_id.eq(user_id)) + .filter(chat_sessions::id.eq(session_id)) + .select(ChatRsSession::as_select()) + .first(self.db) + .await?; + + Ok(session) + } + + pub async fn get_session_with_messages( + &mut self, + user_id: &Uuid, + session_id: &Uuid, + ) -> Result<(ChatRsSession, Vec), diesel::result::Error> { + let (session, messages) = futures::future::try_join( + chat_sessions::table + .filter(chat_sessions::user_id.eq(user_id)) + .filter(chat_sessions::id.eq(session_id)) + .select(ChatRsSession::as_select()) + .first(&mut &**self.db), + chat_messages::table + .filter(chat_messages::session_id.eq(session_id)) + .select(ChatRsMessage::as_select()) + .order_by(chat_messages::created_at.asc()) + .load(&mut &**self.db), + ) + .await?; + + Ok((session, messages)) + } + + pub async fn search_sessions( + &mut self, + user_id: &Uuid, + query: &str, + ) -> Result, diesel::result::Error> { + let sessions = full_text_query(self.db, user_id, query, 10).await?; + + Ok(sessions) + } + + pub async fn update_session( + &mut self, + user_id: &Uuid, + session_id: &Uuid, + data: UpdateChatRsSession<'_>, + ) -> Result { + let updated_id: Uuid = diesel::update(chat_sessions::table.find(session_id)) + .set(data) + .filter(chat_sessions::user_id.eq(user_id)) + .returning(chat_sessions::id) + .get_result(self.db) + .await?; + + Ok(updated_id) + } + + pub async fn delete_session( + &mut self, + user_id: &Uuid, + session_id: &Uuid, + ) -> Result { + let id: Uuid = diesel::delete(chat_sessions::table.find(session_id)) + .filter(chat_sessions::user_id.eq(user_id)) + .returning(chat_sessions::id) + .get_result(self.db) + .await?; + + Ok(id) + } + + pub async fn delete_by_user( + &mut self, + user_id: &Uuid, + ) -> Result, diesel::result::Error> { + let ids: Vec = diesel::delete(chat_sessions::table) + .filter(chat_sessions::user_id.eq(user_id)) + .returning(chat_sessions::id) + .get_results(self.db) + .await?; + + Ok(ids) + } +} diff --git a/server-new/src/lib.rs b/server-new/src/lib.rs index bedeac0..fecbebc 100644 --- a/server-new/src/lib.rs +++ b/server-new/src/lib.rs @@ -7,6 +7,7 @@ mod config; mod db; mod error; mod extractors; +mod llm; mod plugins; mod services; mod state; diff --git a/server-new/src/services/llm/error.rs b/server-new/src/llm/error.rs similarity index 60% rename from server-new/src/services/llm/error.rs rename to server-new/src/llm/error.rs index dbc7f94..61575d7 100644 --- a/server-new/src/services/llm/error.rs +++ b/server-new/src/llm/error.rs @@ -1,3 +1,5 @@ +use crate::services::stream::error::StreamingError; + /// Errors that can occur in an LLM provider request #[derive(Debug, thiserror::Error)] pub enum LlmRequestError { @@ -5,21 +7,17 @@ pub enum LlmRequestError { Provider(String), } -/// Errors that can occur during LLM streaming +/// Errors that can occur in an LLM stream chunk #[derive(Debug, thiserror::Error)] -pub enum LlmStreamError { +pub enum LlmStreamChunkError { #[error("Provider error: {0}")] Provider(String), #[error("Failed to parse event: {0}")] Parsing(#[from] serde_json::Error), #[error("Failed to decode line: {0}")] Decoding(#[from] tokio_util::codec::LinesCodecError), + #[error(transparent)] + Streaming(#[from] StreamingError), #[error("Stream was cancelled")] StreamCancelled, - // #[error("Redis error: {0}")] - // Redis(#[from] fred::error::Error), - // #[error("Tinistream error: {0}")] - // Tinistream(#[from] crate::stream::TiniError), - // #[error("Websocket error: {0}")] - // Websocket(#[from] reqwest_websocket::Error), } diff --git a/server-new/src/services/llm/interface.rs b/server-new/src/llm/interface.rs similarity index 61% rename from server-new/src/services/llm/interface.rs rename to server-new/src/llm/interface.rs index 23d77e7..d053c60 100644 --- a/server-new/src/services/llm/interface.rs +++ b/server-new/src/llm/interface.rs @@ -1,18 +1,21 @@ use futures::{future::BoxFuture, stream::BoxStream}; use super::{ - error::{LlmRequestError, LlmStreamError}, + error::{LlmRequestError, LlmStreamChunkError}, types::{LlmChatRequest, LlmUsage}, }; -/// Trait that all LLM providers must implement -pub trait LlmProvider { - fn stream_chat<'r>(&'r self, request: &'r LlmChatRequest) -> LlmStreamingResponse<'r>; +/// Trait for all LLM providers +pub trait LlmProvider: Send + Sync { + fn stream_chat<'r>(&'r self, request: LlmChatRequest<'r>) -> LlmStreamingResponse<'r>; } +/// Initial API response to a streaming request from the LLM provider pub type LlmStreamingResponse<'r> = BoxFuture<'r, Result>; +/// The response stream from the LLM provider pub type LlmStream = BoxStream<'static, LlmStreamChunkResult>; -pub type LlmStreamChunkResult = Result; +/// The type of the chunks in the LLM response stream +pub type LlmStreamChunkResult = Result; /// A streaming chunk of data from the LLM provider pub enum LlmStreamChunk { diff --git a/server-new/src/services/llm/mod.rs b/server-new/src/llm/mod.rs similarity index 58% rename from server-new/src/services/llm/mod.rs rename to server-new/src/llm/mod.rs index b667ff8..47fe0ca 100644 --- a/server-new/src/services/llm/mod.rs +++ b/server-new/src/llm/mod.rs @@ -1,3 +1,5 @@ +//! LLM interface and provider implementations + pub mod error; pub mod interface; pub mod providers; diff --git a/server-new/src/llm/providers/mod.rs b/server-new/src/llm/providers/mod.rs new file mode 100644 index 0000000..bb2961a --- /dev/null +++ b/server-new/src/llm/providers/mod.rs @@ -0,0 +1,4 @@ +mod openai; +mod utils; + +pub use openai::{OpenAIProvider, OpenAIProviderConfig, OpenAIProviderFlavor}; diff --git a/server-new/src/services/llm/providers/openai/mod.rs b/server-new/src/llm/providers/openai/mod.rs similarity index 58% rename from server-new/src/services/llm/providers/openai/mod.rs rename to server-new/src/llm/providers/openai/mod.rs index cbcdf8e..7c0da58 100644 --- a/server-new/src/services/llm/providers/openai/mod.rs +++ b/server-new/src/llm/providers/openai/mod.rs @@ -2,7 +2,7 @@ use futures::StreamExt; -use crate::services::llm::{ +use crate::llm::{ error::LlmRequestError, interface::{LlmProvider, LlmStreamingResponse}, providers::utils, @@ -17,76 +17,187 @@ use {request::*, response::*}; const OPENAI_API_BASE_URL: &str = "https://api.openai.com/v1"; const OPENROUTER_API_BASE_URL: &str = "https://openrouter.ai/api/v1"; -/// OpenAI chat provider +/// OpenAI-compatible provider behavior variants. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum OpenAIProviderFlavor { + OpenAI, + OpenRouter, +} +impl OpenAIProviderFlavor { + fn name(self) -> &'static str { + match self { + Self::OpenAI => "OpenAI", + Self::OpenRouter => "OpenRouter", + } + } + fn default_base_url(self) -> &'static str { + match self { + Self::OpenAI => OPENAI_API_BASE_URL, + Self::OpenRouter => OPENROUTER_API_BASE_URL, + } + } + fn use_max_completion_tokens(self) -> bool { + self == Self::OpenAI + } + fn include_store_false(self) -> bool { + self == Self::OpenAI + } + fn include_usage_stream_options(self) -> bool { + true + } +} + +/// Configuration for OpenAI-compatible providers. #[derive(Debug, Clone)] -pub struct OpenAIProvider { - client: reqwest::Client, - redis: fred::clients::Client, +pub struct OpenAIProviderConfig { + flavor: OpenAIProviderFlavor, api_key: String, base_url: String, } +impl OpenAIProviderConfig { + pub fn openai(api_key: impl Into) -> Self { + Self::new(OpenAIProviderFlavor::OpenAI, api_key, None::) + } + + pub fn openrouter(api_key: impl Into) -> Self { + Self::new(OpenAIProviderFlavor::OpenRouter, api_key, None::) + } + + pub fn new( + flavor: OpenAIProviderFlavor, + api_key: impl Into, + base_url: Option>, + ) -> Self { + Self { + flavor, + api_key: api_key.into(), + base_url: base_url + .map(Into::into) + .unwrap_or_else(|| flavor.default_base_url().to_owned()) + .trim_end_matches('/') + .to_owned(), + } + } +} + +/// OpenAI-compatible chat provider. +#[derive(Debug, Clone)] +pub struct OpenAIProvider { + client: reqwest::Client, + _redis: fred::clients::Client, + config: OpenAIProviderConfig, +} + impl OpenAIProvider { pub fn new( http_client: &reqwest::Client, redis: &fred::clients::Client, - api_key: &str, - base_url: Option<&str>, + config: OpenAIProviderConfig, ) -> Self { Self { client: http_client.clone(), - redis: redis.clone(), - api_key: api_key.to_owned(), - base_url: base_url.unwrap_or(OPENAI_API_BASE_URL).to_owned(), + _redis: redis.clone(), + config, } } + + pub fn openai( + http_client: &reqwest::Client, + redis: &fred::clients::Client, + api_key: impl Into, + ) -> Self { + Self::new(http_client, redis, OpenAIProviderConfig::openai(api_key)) + } + + pub fn openrouter( + http_client: &reqwest::Client, + redis: &fred::clients::Client, + api_key: impl Into, + ) -> Self { + Self::new( + http_client, + redis, + OpenAIProviderConfig::openrouter(api_key), + ) + } +} + +#[derive(Debug, Clone, Copy)] +struct OpenAIRequestPolicy { + flavor: OpenAIProviderFlavor, +} + +impl OpenAIRequestPolicy { + fn new(flavor: OpenAIProviderFlavor) -> Self { + Self { flavor } + } + + fn max_tokens(self, max_tokens: Option) -> Option { + (!self.flavor.use_max_completion_tokens()) + .then_some(max_tokens) + .flatten() + } + + fn max_completion_tokens(self, max_tokens: Option) -> Option { + self.flavor + .use_max_completion_tokens() + .then_some(max_tokens) + .flatten() + } + + fn store(self) -> Option { + self.flavor.include_store_false().then_some(false) + } + + fn stream_options(self) -> Option { + self.flavor + .include_usage_stream_options() + .then_some(OpenAIStreamOptions { + include_usage: true, + }) + } } impl LlmProvider for OpenAIProvider { - fn stream_chat<'r>(&'r self, req: &'r LlmChatRequest) -> LlmStreamingResponse<'r> { + fn stream_chat<'r>(&'r self, req: LlmChatRequest<'r>) -> LlmStreamingResponse<'r> { + let policy = OpenAIRequestPolicy::new(self.config.flavor); let openai_messages = build_openai_messages(&req.messages); // let openai_tools = tools.as_ref().map(|t| build_openai_tools(t)); // let request = OpenAIRequest { model: &req.options.model, messages: openai_messages, - // OpenAI official API deprecated `max_tokens` for `max_completion_tokens` - max_tokens: match req.options.max_tokens { - Some(max_tokens) if self.base_url != OPENAI_API_BASE_URL => Some(max_tokens), - _ => None, - }, - max_completion_tokens: match req.options.max_tokens { - Some(max_tokens) if self.base_url == OPENAI_API_BASE_URL => Some(max_tokens), - _ => None, - }, + max_tokens: policy.max_tokens(req.options.max_tokens), + max_completion_tokens: policy.max_completion_tokens(req.options.max_tokens), temperature: req.options.temperature, - store: (self.base_url == OPENAI_API_BASE_URL).then_some(false), + store: policy.store(), stream: Some(true), - stream_options: Some(OpenAIStreamOptions { - include_usage: true, - }), + stream_options: policy.stream_options(), // tools: openai_tools, // modalities: options.modalities.as_ref(), ..Default::default() }; + let provider_name = self.config.flavor.name(); Box::pin(async move { let response = self .client - .post(format!("{}/chat/completions", self.base_url)) - .header("authorization", format!("Bearer {}", self.api_key)) + .post(format!("{}/chat/completions", self.config.base_url)) + .header("authorization", format!("Bearer {}", self.config.api_key)) .header("content-type", "application/json") .json(&request) .send() .await - .map_err(|e| LlmRequestError::Provider(format!("OpenAI request failed: {}", e)))?; + .map_err(|e| { + LlmRequestError::Provider(format!("{provider_name} request failed: {e}")) + })?; if !response.status().is_success() { let status = response.status(); let error_text = response.text().await.unwrap_or_default(); return Err(LlmRequestError::Provider(format!( - "OpenAI API error {}: {}", - status, error_text + "{provider_name} API error {status}: {error_text}", ))); } diff --git a/server-new/src/services/llm/providers/openai/request.rs b/server-new/src/llm/providers/openai/request.rs similarity index 99% rename from server-new/src/services/llm/providers/openai/request.rs rename to server-new/src/llm/providers/openai/request.rs index 8ef0256..afa8043 100644 --- a/server-new/src/services/llm/providers/openai/request.rs +++ b/server-new/src/llm/providers/openai/request.rs @@ -1,6 +1,6 @@ use serde::Serialize; -use crate::services::llm::{ +use crate::llm::{ providers::utils, types::{LlmFileType, LlmMessage}, }; diff --git a/server-new/src/services/llm/providers/openai/response.rs b/server-new/src/llm/providers/openai/response.rs similarity index 96% rename from server-new/src/services/llm/providers/openai/response.rs rename to server-new/src/llm/providers/openai/response.rs index 0bf881f..ab26f5c 100644 --- a/server-new/src/services/llm/providers/openai/response.rs +++ b/server-new/src/llm/providers/openai/response.rs @@ -1,12 +1,12 @@ use serde::Deserialize; -use crate::services::llm::{error::LlmStreamError, interface::LlmStreamChunk, types::LlmUsage}; +use crate::llm::{error::LlmStreamChunkError, interface::LlmStreamChunk, types::LlmUsage}; /// Parse chunks from an OpenAI SSE event pub fn parse_openai_event( mut event: OpenAIStreamResponse, - tool_calls: &mut Vec, -) -> Vec> { + _tool_calls: &mut Vec, +) -> Vec> { let mut chunks = Vec::with_capacity(1); if let Some(delta) = event.choices.pop().and_then(|c| c.delta) { if let Some(text) = delta.content { diff --git a/server-new/src/services/llm/providers/utils.rs b/server-new/src/llm/providers/utils.rs similarity index 84% rename from server-new/src/services/llm/providers/utils.rs rename to server-new/src/llm/providers/utils.rs index 62ed993..327f8f7 100644 --- a/server-new/src/services/llm/providers/utils.rs +++ b/server-new/src/llm/providers/utils.rs @@ -8,7 +8,7 @@ use tokio_util::{ io::StreamReader, }; -use crate::services::llm::error::LlmStreamError; +use crate::llm::error::LlmStreamChunkError; /// Create a data URI pub fn create_data_uri(content_type: &str, b64_string: &str) -> String { @@ -18,7 +18,7 @@ pub fn create_data_uri(content_type: &str, b64_string: &str) -> String { /// Get a stream of deserialized events from a provider SSE stream. pub fn get_sse_events( response: reqwest::Response, -) -> impl Stream> { +) -> impl Stream> { let stream_reader = StreamReader::new(response.bytes_stream().map_err(std::io::Error::other)); let line_reader = FramedRead::new(stream_reader, LinesCodec::new()); @@ -30,13 +30,13 @@ pub fn get_sse_events( if data.trim_start().is_empty() || data == "[DONE]" { None // Skip empty lines and termination markers } else { - Some(serde_json::from_str::(data).map_err(LlmStreamError::Parsing)) + Some(serde_json::from_str::(data).map_err(LlmStreamChunkError::Parsing)) } } else { None // Ignore non-data lines } } - Err(e) => Some(Err(LlmStreamError::Decoding(e))), + Err(e) => Some(Err(LlmStreamChunkError::Decoding(e))), } }) } @@ -44,11 +44,11 @@ pub fn get_sse_events( /// Get a stream of deserialized events from a provider JSON stream, not SSE (e.g. Ollama uses this format). pub fn get_json_events( response: reqwest::Response, -) -> impl Stream> { +) -> impl Stream> { let stream_reader = StreamReader::new(response.bytes_stream().map_err(std::io::Error::other)); let line_reader = FramedRead::new(stream_reader, LinesCodec::new()); line_reader.map(|line_result| match line_result { - Ok(line) => serde_json::from_str::(&line).map_err(LlmStreamError::Parsing), - Err(e) => Err(LlmStreamError::Decoding(e)), + Ok(line) => serde_json::from_str::(&line).map_err(LlmStreamChunkError::Parsing), + Err(e) => Err(LlmStreamChunkError::Decoding(e)), }) } diff --git a/server-new/src/services/llm/types.rs b/server-new/src/llm/types.rs similarity index 95% rename from server-new/src/services/llm/types.rs rename to server-new/src/llm/types.rs index 9603a3f..9075b06 100644 --- a/server-new/src/services/llm/types.rs +++ b/server-new/src/llm/types.rs @@ -1,10 +1,10 @@ use serde::{Deserialize, Serialize}; /// Generic chat request for all LLM providers -pub struct LlmChatRequest { - pub messages: Vec, +pub struct LlmChatRequest<'r> { + pub messages: &'r [LlmMessage], // tools: Option>, - pub options: LlmChatOptions, + pub options: &'r LlmChatOptions, } /// Generic message type to send to LLM providers diff --git a/server-new/src/services/chat/error.rs b/server-new/src/services/chat/error.rs index c3b814b..dd45481 100644 --- a/server-new/src/services/chat/error.rs +++ b/server-new/src/services/chat/error.rs @@ -1,8 +1,14 @@ -use crate::services::llm::error::LlmRequestError; - /// Chat service errors #[derive(Debug, thiserror::Error)] pub enum ChatError { #[error(transparent)] - Request(#[from] LlmRequestError), + Request(#[from] crate::llm::error::LlmRequestError), + #[error("Invalid message history")] + Messages, + #[error(transparent)] + Streaming(#[from] crate::services::stream::error::StreamingError), + #[error("database error: {0}")] + Database(#[from] diesel::result::Error), + #[error("database pool error: {0}")] + DatabasePool(#[from] crate::db::DbPoolError), } diff --git a/server-new/src/services/chat/messages.rs b/server-new/src/services/chat/messages.rs new file mode 100644 index 0000000..bac8902 --- /dev/null +++ b/server-new/src/services/chat/messages.rs @@ -0,0 +1,76 @@ +use crate::{ + db::models::{ChatRsMessage, ChatRsMessageRole}, + llm::types::{LlmAssistantMessage, LlmMessage, LlmUserMessage}, + services::chat::error::ChatError, +}; + +/// Extract any attached files, then convert the database messages to the generic format +/// for sending to LLM providers +pub fn build_llm_messages( + messages: Vec, + // user_id: &Uuid, + // session_id: &Uuid, + // db: &mut DbConnection, + // storage: &LocalStorage, +) -> Result, ChatError> { + // // Get content of any attached files in the messages + // let mut file_map: HashMap = HashMap::new(); + // let file_ids: Vec = messages.iter().fold(Vec::new(), |mut acc, message| { + // if let Some(file_ids) = message.meta.user.as_ref().and_then(|u| u.files.as_ref()) { + // acc.extend(file_ids); + // } + // acc + // }); + // for file_id in file_ids { + // let file = FileDbService::new(db) + // .find_session_file(user_id, session_id, &file_id) + // .await?; + // let (file_type, content) = file.read_to_string(Some(session_id), storage).await?; + // file_map.insert( + // file_id, + // LlmFileInput { + // name: file.path, + // content_type: file.content_type, + // file_type, + // content, + // }, + // ); + // } + + // Convert the messages + let llm_messages = messages + .into_iter() + .map(|message| match message.role { + ChatRsMessageRole::User => { + // let files = message.meta.user.and_then(|u| u.files).map(|file_ids| { + // file_ids + // .iter() + // .filter_map(|id| file_map.remove(id)) + // .collect() + // }); + Ok(LlmMessage::User(LlmUserMessage { + text: message.content, + files: None, + })) + } + ChatRsMessageRole::Assistant => Ok(LlmMessage::Assistant(LlmAssistantMessage { + text: message.content, + // tool_calls: message.meta.assistant.and_then(|a| a.tool_calls), + })), + ChatRsMessageRole::System => Ok(LlmMessage::System(message.content)), + ChatRsMessageRole::Tool => { + // if let Some(tool_call) = message.meta.tool_call { + // Ok(LlmMessage::Tool(LlmToolResult { + // tool_call_id: tool_call.id, + // tool_name: tool_call.tool_name, + // content: message.content, + // })) + // } else { + Err(ChatError::Messages) + // } + } + }) + .collect::, ChatError>>()?; + + Ok(llm_messages) +} diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index 3a04855..cff77d7 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -1,43 +1,133 @@ -use futures::Stream; +use tinistream_client::types::StreamAccessResponse; +use uuid::Uuid; -use crate::services::{ - chat::error::ChatError, +use crate::{ + db::{ + DbPool, DbService, + models::{ + AssistantMeta, ChatRsMessage, ChatRsMessageMeta, ChatRsMessageRole, NewChatRsMessage, + UserMeta, + }, + }, llm::{ - interface::{LlmProvider, LlmStreamChunkResult}, - providers::OpenAIProvider, - types::{LlmChatOptions, LlmChatRequest, LlmMessage, LlmUserMessage}, + interface::LlmProvider, + types::{LlmChatOptions, LlmChatRequest, LlmUserMessage}, + }, + services::{ + chat::error::ChatError, + stream::{LlmStreamOutput, StreamingService, tinistream::TinistreamClient}, }, }; mod error; +mod messages; -pub struct ChatService<'a> { - http_client: &'a reqwest::Client, - redis: &'a fred::clients::Client, +pub struct ChatService<'r> { + db_pool: &'r DbPool, + tinistream: &'r TinistreamClient, } -impl<'a> ChatService<'a> { - pub fn new(http_client: &'a reqwest::Client, redis: &'a fred::clients::Client) -> Self { - Self { http_client, redis } +impl<'r> ChatService<'r> { + pub fn new(db_pool: &'r DbPool, tinistream: &'r TinistreamClient) -> Self { + Self { + db_pool, + tinistream, + } } - pub async fn test_chat(&self) -> Result, ChatError> { - let provider: Box = - Box::new(OpenAIProvider::new(self.http_client, self.redis, "", None)); - let messages = vec![LlmMessage::User(LlmUserMessage { - text: "Hello!".into(), + pub async fn stream_user_chat( + &self, + user_id: Uuid, + session_id: Uuid, + provider: &dyn LlmProvider, + provider_id: i32, + user_message: LlmUserMessage, + chat_options: LlmChatOptions, + ) -> Result { + // Check for existing chat session, then save the new user message to it + let mut db = DbService::from_pool(self.db_pool).await?; + let (_existing_session, mut session_messages) = db + .chats() + .get_session_with_messages(&user_id, &session_id) + .await?; + let new_message = db + .chats() + .save_message(NewChatRsMessage { + session_id: &session_id, + role: ChatRsMessageRole::User, + content: &user_message.text, + meta: ChatRsMessageMeta::new_user(UserMeta::default()), + }) + .await?; + session_messages.push(new_message); + + // Send the request to the LLM provider and get the streaming response + let llm_messages = messages::build_llm_messages(session_messages)?; + let stream = provider + .stream_chat(LlmChatRequest { + messages: &llm_messages, + options: &chat_options, + }) + .await?; + + // Create a new client stream in `tinistream` to stream the response to the user + let stream_key = StreamingService::chat_stream_key(&user_id, &session_id); + let (stream_access, ws_writer, ws_reader) = StreamingService::new(self.tinistream) + .create_stream(&stream_key) + .await?; + + // Spawn task to process and save the streaming response + let db_pool = self.db_pool.to_owned(); + let tinistream_client = self.tinistream.to_owned(); + tokio::spawn(async move { + let response = StreamingService::process_stream(stream, ws_writer, ws_reader).await; + let stream_cancelled = response.cancelled; + if let Err(err) = + Self::persist_response(db_pool, &session_id, provider_id, chat_options, response) + .await + { + tracing::error!("Failed to save assistant response: {err}"); + } + if !stream_cancelled { + let _ = StreamingService::new(&tinistream_client) + .end_stream(&stream_key) + .await; + } + }); + + // Return the URL and token for the user to access the client stream + Ok(stream_access) + } + + /// Save response message and metadata to database + async fn persist_response( + db_pool: DbPool, + session_id: &Uuid, + provider_id: i32, + chat_options: LlmChatOptions, + response: LlmStreamOutput, + ) -> Result { + let mut db = DbService::from_pool(&db_pool).await?; + let assistant_meta = AssistantMeta { + provider_id, + provider_options: Some(chat_options), + // tool_calls: response.tool_calls, + // files: image_ids, + usage: response.usage, + errors: response.errors, + partial: response.cancelled.then_some(true), ..Default::default() - })]; - let request = LlmChatRequest { - messages: messages, - options: LlmChatOptions { - model: "gpt-5-mini".into(), - ..Default::default() - }, }; + let new_message = db + .chats() + .save_message(NewChatRsMessage { + content: &response.text.unwrap_or_default(), + meta: ChatRsMessageMeta::new_assistant(assistant_meta), + role: ChatRsMessageRole::Assistant, + session_id: &session_id, + }) + .await?; - let response = provider.stream_chat(&request).await?; - - Ok(response) + Ok(new_message) } } diff --git a/server-new/src/services/llm/providers/mod.rs b/server-new/src/services/llm/providers/mod.rs deleted file mode 100644 index c219eef..0000000 --- a/server-new/src/services/llm/providers/mod.rs +++ /dev/null @@ -1,4 +0,0 @@ -mod openai; -mod utils; - -pub use openai::OpenAIProvider; diff --git a/server-new/src/services/mod.rs b/server-new/src/services/mod.rs index 0833d30..d2d9b1b 100644 --- a/server-new/src/services/mod.rs +++ b/server-new/src/services/mod.rs @@ -1,3 +1,3 @@ pub mod auth; pub mod chat; -pub mod llm; +pub mod stream; diff --git a/server-new/src/services/stream/error.rs b/server-new/src/services/stream/error.rs new file mode 100644 index 0000000..fb41cc5 --- /dev/null +++ b/server-new/src/services/stream/error.rs @@ -0,0 +1,8 @@ +/// Errors that can occur during streaming +#[derive(Debug, thiserror::Error)] +pub enum StreamingError { + #[error("Client streaming error: {0}")] + Tinistream(#[from] super::tinistream::TiniError), + #[error("Websocket error: {0}")] + Websocket(#[from] reqwest_websocket::Error), +} diff --git a/server-new/src/services/stream/mod.rs b/server-new/src/services/stream/mod.rs new file mode 100644 index 0000000..b4d4324 --- /dev/null +++ b/server-new/src/services/stream/mod.rs @@ -0,0 +1,82 @@ +use futures::{ + StreamExt, + stream::{SplitSink, SplitStream}, +}; +use reqwest_websocket::WebSocket; +use tinistream::TinistreamClient; +use tinistream_client::types::{StreamAccessResponse, StreamStatus}; +use uuid::Uuid; + +pub mod error; +pub mod tinistream; +mod writer; + +#[cfg(test)] +mod tests; + +use crate::{ + llm::{interface::LlmStream, types::LlmUsage}, + services::stream::error::StreamingError, +}; + +/// Handles stream processing and interacting with `tinistream` for streaming to users +pub struct StreamingService<'r> { + tinistream: &'r TinistreamClient, +} + +/// Complete, accumulated response from the LLM provider stream +pub struct LlmStreamOutput { + pub text: Option, + // pub tool_calls: Option>, + // pub images: Option>, + pub usage: Option, + pub errors: Option>, + pub cancelled: bool, +} + +pub type WsWriter = SplitSink; +pub type WsReader = SplitStream; + +impl<'r> StreamingService<'r> { + pub fn new(tinistream: &'r TinistreamClient) -> Self { + Self { tinistream } + } + + /// Get the key of the chat stream in Redis for the given user and session ID + pub fn chat_stream_key(user_id: &Uuid, session_id: &Uuid) -> String { + format!("{}{}", Self::chat_stream_prefix(user_id), session_id) + } + + /// Get the key prefix for the user's chat streams in Redis + pub fn chat_stream_prefix(user_id: &Uuid) -> String { + format!("user:{}:chat:", user_id) + } + + /// Start the client stream in `tinistream`, and return a WebSocket writer and reader for it + pub async fn create_stream( + &self, + stream_key: &str, + ) -> Result<(StreamAccessResponse, WsWriter, WsReader), StreamingError> { + let stream_access = self.tinistream.stream_start(stream_key).await?; + let (writer, reader) = self.tinistream.stream_writer_ws(stream_key).await?.split(); + + Ok((stream_access, writer, reader)) + } + + /// Process and write the LLM stream to `tinistream` via the WebSocket connection, + /// and return the accumulated response. + pub async fn process_stream( + stream: LlmStream, + writer: WsWriter, + reader: WsReader, + ) -> LlmStreamOutput { + writer::LlmStreamWriter::new() + .process(stream, writer, reader) + .await + } + + /// Signal end of stream in `tinistream` + pub async fn end_stream(&self, stream_key: &str) -> Result { + Ok(self.tinistream.stream_end(stream_key).await?) + } +} diff --git a/server-new/src/services/stream/tests/mod.rs b/server-new/src/services/stream/tests/mod.rs new file mode 100644 index 0000000..2046c4e --- /dev/null +++ b/server-new/src/services/stream/tests/mod.rs @@ -0,0 +1,252 @@ +// use futures::StreamExt; +// use reqwest_websocket::WebSocket; +// use uuid::Uuid; + +// use crate::services::stream::{StreamService, writer::LlmStreamWriter}; + +// mod utils; +// use utils::*; + +// async fn create_test_writer( +// user_id: &Uuid, +// session_id: &Uuid, +// ) -> (String, WebSocket, LlmStreamWriter) { +// let key = StreamService::chat_stream_key(user_id, session_id); +// let tini = setup_tini_client(); +// tini.stream_start(&key).await.expect("should start stream"); +// let ws = tini +// .stream_writer_ws(&key) +// .await +// .expect("should connect to WebSocket for adding events"); + +// (key, ws, LlmStreamWriter::new()) +// } + +// #[tokio::test] +// async fn stream_writer_basic_functionality() { +// let user_id = Uuid::new_v4(); +// let session_id = Uuid::new_v4(); +// let tini = setup_tini_client(); +// let (key, ws, mut writer) = create_test_writer(&user_id, &session_id).await; +// let (ws_writer, ws_reader) = ws.split(); + +// // Create stream +// assert!(tini.stream_exists(&key).await.unwrap()); + +// // Create Lorem provider and get stream +// let lorem = LoremProvider::new(); +// let stream = lorem +// .chat_stream(vec![], None, &LlmProviderOptions::default()) +// .await +// .expect("Failed to create lorem stream"); + +// // Process the stream +// let LlmOutput { +// text, +// tool_calls, +// usage, +// errors, +// cancelled, +// .. +// } = writer.process(stream, ws_writer, ws_reader).await; + +// // Verify results +// assert!(text.is_some()); +// let text = text.unwrap(); +// assert!(!text.is_empty()); +// assert!(text.contains("Lorem ipsum")); +// assert!(text.contains("dolor sit")); + +// assert!(tool_calls.is_none()); +// assert!(usage.is_none()); +// assert!(errors.is_some()); // Lorem provider generates some test errors +// assert!(!cancelled); + +// // End stream +// assert!(tini.stream_end(&key).await.is_ok()); + +// // Stream should be deleted after end +// assert!(!tini.stream_exists(&key).await.unwrap()); +// } + +// #[tokio::test] +// async fn stream_writer_batching() { +// let user_id = Uuid::new_v4(); +// let session_id = Uuid::new_v4(); +// let tini = setup_tini_client(); +// let (key, ws, mut writer) = create_test_writer(&user_id, &session_id).await; +// let (ws_writer, ws_reader) = ws.split(); + +// // Create a custom stream with small chunks to test batching +// let chunks = vec![ +// "Hello", " ", "world", "!", " ", "This", " ", "is", " ", "a", " ", "test", +// ]; +// let chunk_stream = tokio_stream::iter( +// chunks +// .into_iter() +// .map(|text| Ok(LlmStreamChunk::Text(text.into()))), +// ); + +// let stream: LlmStream = Box::pin(chunk_stream); +// let LlmOutput { +// text, cancelled, .. +// } = writer.process(stream, ws_writer, ws_reader).await; + +// assert!(text.is_some()); +// let text = text.unwrap(); +// assert_eq!(text, "Hello world! This is a test"); +// assert!(!cancelled); + +// tini.stream_end(&key).await.ok(); +// } + +// #[tokio::test] +// async fn stream_writer_error_handling() { +// let user_id = Uuid::new_v4(); +// let session_id = Uuid::new_v4(); +// let tini = setup_tini_client(); +// let (key, ws, mut writer) = create_test_writer(&user_id, &session_id).await; +// let (ws_writer, ws_reader) = ws.split(); + +// // Create a stream that produces an error +// let error_stream = tokio_stream::iter(vec![ +// Ok(LlmStreamChunk::Text("Hello".to_string())), +// Err(LlmStreamError::ProviderError("Test error".into())), +// Ok(LlmStreamChunk::Text(" World".to_string())), +// ]); + +// let stream: LlmStream = Box::pin(error_stream); +// let LlmOutput { +// text, +// errors, +// cancelled, +// .. +// } = writer.process(stream, ws_writer, ws_reader).await; + +// assert!(text.is_some()); +// let text = text.unwrap(); +// assert_eq!(text, "Hello World"); + +// assert!(errors.is_some()); +// let errors = errors.unwrap(); +// assert!(!errors.is_empty()); +// assert!(errors.iter().any(|e| e.contains("Test error"))); + +// assert!(!cancelled); + +// tini.stream_end(&key).await.ok(); +// } + +// #[tokio::test] +// async fn stream_writer_cancel() { +// let user_id = Uuid::new_v4(); +// let session_id = Uuid::new_v4(); +// let tini = setup_tini_client(); +// let (key, ws, mut writer) = create_test_writer(&user_id, &session_id).await; +// let (ws_writer, ws_reader) = ws.split(); + +// assert!(tini.stream_exists(&key).await.unwrap()); + +// let stream = LoremProvider::new() +// .chat_stream(vec![], None, &LlmProviderOptions::default()) +// .await +// .expect("Failed to create lorem stream"); +// let process_fut = writer.process(stream, ws_writer, ws_reader); + +// // Cancel the stream after 2 seconds +// tokio::time::sleep(std::time::Duration::from_secs(2)).await; +// tini.stream_cancel(&key).await.unwrap(); + +// // process() response should show that stream was cancelled +// let LlmOutput { +// errors, cancelled, .. +// } = process_fut.await; +// assert!(cancelled); +// assert!(errors.unwrap().last().unwrap().contains("cancelled")); + +// // Stream should be deleted after cancel +// assert!(!tini.stream_exists(&key).await.unwrap()); +// } + +// #[tokio::test] +// async fn stream_writer_usage_tracking() { +// let user_id = Uuid::new_v4(); +// let session_id = Uuid::new_v4(); +// let tini = setup_tini_client(); +// let (key, ws, mut writer) = create_test_writer(&user_id, &session_id).await; +// let (ws_writer, ws_reader) = ws.split(); + +// assert!(tini.stream_exists(&key).await.unwrap()); + +// // Create a stream with usage information +// let usage_stream = tokio_stream::iter(vec![ +// Ok(LlmStreamChunk::Text("Hello".into())), +// Ok(LlmStreamChunk::Usage(LlmUsage { +// input_tokens: Some(10), +// output_tokens: Some(5), +// cost: Some(0.001), +// })), +// Ok(LlmStreamChunk::Text(" World".into())), +// Ok(LlmStreamChunk::Usage(LlmUsage { +// input_tokens: None, // Should not override +// output_tokens: Some(7), // Should update +// cost: Some(0.002), // Should update +// })), +// ]); + +// let stream: LlmStream = Box::pin(usage_stream); +// let LlmOutput { +// text, +// usage, +// cancelled, +// .. +// } = writer.process(stream, ws_writer, ws_reader).await; + +// assert!(text.is_some()); +// assert_eq!(text.unwrap(), "Hello World"); + +// assert!(usage.is_some()); +// let usage = usage.unwrap(); +// assert_eq!(usage.input_tokens, Some(10)); +// assert_eq!(usage.output_tokens, Some(7)); +// assert_eq!(usage.cost, Some(0.002)); + +// assert!(!cancelled); + +// tini.stream_end(&key).await.ok(); +// } + +// #[tokio::test] +// async fn redis_stream_entries() { +// let user_id = Uuid::new_v4(); +// let session_id = Uuid::new_v4(); +// let tini = setup_tini_client(); +// let (key, ws, mut writer) = create_test_writer(&user_id, &session_id).await; +// let (ws_writer, ws_reader) = ws.split(); + +// assert!(tini.stream_exists(&key).await.unwrap()); + +// // Verify start event was written +// let info = tini +// .stream_info(&key) +// .await +// .expect("Failed to check stream") +// .expect("Stream not found"); +// assert_eq!(info.length, 1); + +// // Create a simple stream +// let stream = tokio_stream::iter(vec![Ok(LlmStreamChunk::Text("Test chunk".into()))]).boxed(); +// writer.process(stream, ws_writer, ws_reader).await; +// drop(writer); + +// // Should have start + text entries +// tokio::time::sleep(std::time::Duration::from_secs(1)).await; +// let info = tini +// .stream_info(&key) +// .await +// .expect("Failed to check stream") +// .expect("Stream not found"); +// assert_eq!(info.length, 2); + +// tini.stream_end(&key).await.ok(); +// } diff --git a/server-new/src/services/stream/tests/utils.rs b/server-new/src/services/stream/tests/utils.rs new file mode 100644 index 0000000..3fa89b4 --- /dev/null +++ b/server-new/src/services/stream/tests/utils.rs @@ -0,0 +1,15 @@ +use super::super::tinistream::TinistreamClient; + +pub fn setup_tini_client() -> TinistreamClient { + let url = dotenvy::var("RS_CHAT_TINISTREAM_URL").unwrap_or("http://127.0.0.1:8081".to_owned()); + let api_key = dotenvy::var("RS_CHAT_TINISTREAM_API_KEY").unwrap_or("".to_owned()); + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert("X-API-KEY", api_key.parse().expect("Should be valid")); + let tini_http_client = reqwest::ClientBuilder::new() + .default_headers(headers) + .build() + .expect("Failed to build tinistream HTTP client"); + + let tini_client = tinistream_client::Client::new_with_client(&url, tini_http_client.clone()); + TinistreamClient::new(tini_client) +} diff --git a/server-new/src/services/stream/tinistream.rs b/server-new/src/services/stream/tinistream.rs new file mode 100644 index 0000000..5000ca6 --- /dev/null +++ b/server-new/src/services/stream/tinistream.rs @@ -0,0 +1,161 @@ +//! Client for `tinistream` to handle streaming responses + +use reqwest_websocket::{Upgrade, WebSocket}; +use tinistream_client::{Client, ClientEventsExt, ClientInfo, ClientStreamExt, Error, types::*}; + +/// A client for interacting with the `tinistream` API. +#[derive(Debug, Clone)] +pub struct TinistreamClient { + client: Client, +} + +/// Result type for tinistream API operations. +pub type TiniResult = Result; + +#[derive(Debug, thiserror::Error)] +#[error("{message}")] +pub struct TiniError { + pub status: u16, + pub code: String, + pub message: String, +} + +impl TinistreamClient { + pub fn new(client: Client) -> Self { + Self { client } + } + + /// Returns a list of keys with the given prefix that have an active stream. + pub async fn active_streams(&self, prefix: &str) -> TiniResult> { + let streams = self + .client + .list_streams() + .pattern(format!("{prefix}*")) + .send() + .await? + .into_inner(); + Ok(streams) + } + + /// Returns whether an active chat stream exists for the given key. + pub async fn stream_exists(&self, key: &str) -> TiniResult { + match self.client.get_stream_info().key(key).send().await { + Ok(_) => Ok(true), + Err(err) => match err.status() { + Some(reqwest::StatusCode::NOT_FOUND) => Ok(false), + _ => Err(err.into()), + }, + } + } + + /// Returns info about a chat stream at the given key. + pub async fn stream_info(&self, key: &str) -> TiniResult> { + let info = self.client.get_stream_info().key(key).send().await?; + Ok(Some(info.into_inner())) + } + + /// Start the chat stream and get the client URL and access token + pub async fn stream_start(&self, key: &str) -> TiniResult { + let res = self + .client + .create_stream() + .body(StreamRequest::builder().key(key)) + .send() + .await?; + Ok(res.into_inner()) + } + + /// Get URL and token for a client to access a stream + pub async fn stream_connect(&self, key: &str) -> TiniResult { + let res = self + .client + .create_token() + .body(StreamRequest::builder().key(key)) + .send() + .await?; + Ok(res.into_inner()) + } + + /// Get a WebSocket connection to write to a stream + pub async fn stream_writer_ws(&self, key: &str) -> Result { + let http_client = self.client.client(); + let res = http_client + .get(format!("{}/api/event/add/ws-stream", self.client.baseurl())) + .query(&[("key", key)]) + .upgrade() + .send() + .await?; + res.into_websocket().await + } + + pub async fn stream_add( + &self, + key: &str, + events: Vec, + ) -> TiniResult> { + let events = events + .into_iter() + .map(|event| event.try_into()) + .collect::, _>>()?; + let res = self + .client + .add_events() + .body(AddEventsRequest::builder().key(key).events(events)) + .send() + .await?; + Ok(res.into_inner().ids) + } + + /// Cancel a stream + pub async fn stream_cancel(&self, key: &str) -> TiniResult { + let res = self + .client + .cancel_stream() + .body(StreamRequest::builder().key(key)) + .send() + .await?; + Ok(res.into_inner().status) + } + + /// Signal the end of a stream + pub async fn stream_end(&self, key: &str) -> TiniResult { + let res = self + .client + .end_stream() + .body(StreamRequest::builder().key(key)) + .send() + .await?; + Ok(res.into_inner().status) + } +} + +impl From> for TiniError { + fn from(value: Error) -> Self { + match value { + Error::ErrorResponse(res) => { + let status = res.status().as_u16(); + let res = res.into_inner(); + TiniError { + status, + code: res.code, + message: res.message, + } + } + res => TiniError { + status: res.status().map_or(500, |s| s.as_u16()), + code: "unexpected".to_owned(), + message: res.to_string(), + }, + } + } +} + +impl From for TiniError { + fn from(value: error::ConversionError) -> Self { + TiniError { + status: 400, + code: "invalid_event".to_owned(), + message: value.to_string(), + } + } +} diff --git a/server-new/src/services/stream/writer.rs b/server-new/src/services/stream/writer.rs new file mode 100644 index 0000000..bde6c77 --- /dev/null +++ b/server-new/src/services/stream/writer.rs @@ -0,0 +1,250 @@ +use std::time::{Duration, Instant}; + +use futures::{SinkExt, StreamExt}; +use reqwest_websocket::Message as WsMessage; +use serde::Serialize; +use tokio_util::sync::CancellationToken; + +use crate::{ + llm::{ + error::LlmStreamChunkError, + interface::{LlmStream, LlmStreamChunk}, + types::LlmUsage, + }, + services::stream::{LlmStreamOutput, WsReader, WsWriter, error::StreamingError}, +}; + +/// Interval at which chunks are flushed to the Redis stream. +const FLUSH_INTERVAL: Duration = Duration::from_millis(400); +/// Max # of characters of the text chunk before it is automatically flushed to Redis. +const MAX_CHUNK_SIZE: usize = 75; + +/// Utility for processing an incoming LLM response stream and writing chunks to `tinistream`. +#[derive(Debug)] +pub struct LlmStreamWriter { + /// The current chunk of data being processed. + current_chunk: ChunkState, + /// Accumulated text response from the assistant. + complete_text: Option, + /// Accumulated tool calls from the assistant. + // tool_calls: Option>, + /// Accumulated generated images from the assistant. + // images: Option>, + /// Accumulated errors during the stream from the LLM provider. + errors: Option>, + /// Accumulated usage information from the LLM provider. + usage: Option, +} + +/// Internal state +#[derive(Debug, Default)] +struct ChunkState { + text: Option, + // tool_calls: Option>, + // pending_tool_calls: Option>, + error: Option, +} + +/// Chunk of the LLM response stored in the Redis stream. +#[derive(Debug, Serialize)] +#[serde(tag = "event", content = "data", rename_all = "snake_case")] +pub(super) enum RedisStreamChunk { + Text(String), + ToolCall(String), + PendingToolCall(String), + Error(String), +} + +impl LlmStreamWriter { + pub fn new() -> Self { + LlmStreamWriter { + current_chunk: ChunkState::default(), + complete_text: None, + // tool_calls: None, + // images: None, + errors: None, + usage: None, + } + } + + /// Process the incoming stream from the LLM provider, intermittently flushing + /// chunks to `tinistream` via the WebSocket connection, and return the final + /// accumulated response. + pub async fn process( + &mut self, + stream: LlmStream, + mut writer: WsWriter, + mut reader: WsReader, + ) -> LlmStreamOutput { + let mut cancelled = false; + + // Spawn task to listen for stream cancellation + let cancel_token = CancellationToken::new(); + let cancel_task_token = cancel_token.clone(); + let cancel_task = tokio::spawn(async move { + while let Some(res) = reader.next().await { + if let Ok(WsMessage::Close { .. }) = res { + cancel_task_token.cancel(); + } + } + }); + + tokio::select! { + _ = self.process_stream(stream, &mut writer) => {} + _ = cancel_token.cancelled() => { + self.errors.get_or_insert_default().push(LlmStreamChunkError::StreamCancelled); + cancelled = true; + } + } + + cancel_task.abort(); + writer.close().await.ok(); + + LlmStreamOutput { + text: self.complete_text.take(), + // tool_calls: self.tool_calls.take(), + // images: self.images.take(), + usage: self.usage.take(), + errors: self.errors.take().map(|e| { + e.into_iter() + .map(|e| e.to_string()) + .collect::>() + }), + cancelled, + } + } + + async fn process_stream(&mut self, mut stream: LlmStream, writer: &mut WsWriter) { + let mut last_flush_time = Instant::now(); + loop { + match stream.next().await { + Some(Ok(chunk)) => match chunk { + LlmStreamChunk::Text(text) => self.process_text(&text), + // LlmStreamChunk::ToolCalls(tool_calls) => self.process_tool_calls(tool_calls), + // LlmStreamChunk::PendingToolCall(pending_tool_call) => { + // self.process_pending_tool_call(pending_tool_call) + // } + // LlmStreamChunk::Images(images) => self.process_images(images), + LlmStreamChunk::Usage(usage) => self.process_usage(usage), + }, + Some(Err(err)) => self.process_error(err), + None => break, + } + + if self.should_flush(&last_flush_time) { + if let Err(err) = self.flush_chunks(writer).await { + self.process_error(LlmStreamChunkError::from(err)); + } + last_flush_time = Instant::now(); + } + } + + if let Err(err) = self.flush_chunks(writer).await { + self.process_error(LlmStreamChunkError::from(err)); + } + } + + fn process_text(&mut self, text: &str) { + self.current_chunk + .text + .get_or_insert_with(|| String::with_capacity(MAX_CHUNK_SIZE)) + .push_str(text); + self.complete_text + .get_or_insert_with(|| String::with_capacity(1024)) + .push_str(text); + } + + // fn process_tool_calls(&mut self, tool_calls: Vec) { + // self.current_chunk + // .tool_calls + // .get_or_insert_default() + // .extend(tool_calls.clone()); + // self.tool_calls.get_or_insert_default().extend(tool_calls); + // } + + // fn process_pending_tool_call(&mut self, tool_call: LlmPendingToolCall) { + // let current_chunk = self + // .current_chunk + // .pending_tool_calls + // .get_or_insert_default(); + // if !current_chunk.iter().any(|tc| tc.index == tool_call.index) { + // current_chunk.push(tool_call); + // } + // } + + // fn process_images(&mut self, images: Vec) { + // self.images.get_or_insert_default().extend(images); + // } + + fn process_usage(&mut self, usage_chunk: LlmUsage) { + let usage = self.usage.get_or_insert_default(); + if let Some(input_tokens) = usage_chunk.input_tokens { + usage.input_tokens = Some(input_tokens); + } + if let Some(output_tokens) = usage_chunk.output_tokens { + usage.output_tokens = Some(output_tokens); + } + if let Some(cost) = usage_chunk.cost { + usage.cost = Some(cost); + } + } + + fn process_error(&mut self, err: LlmStreamChunkError) { + self.current_chunk.error = Some(err.to_string()); + self.errors.get_or_insert_default().push(err); + } + + fn should_flush(&self, last_flush_time: &Instant) -> bool { + // if self.current_chunk.tool_calls.is_some() || self.current_chunk.error.is_some() { + // return true; + // } + if self.current_chunk.error.is_some() { + return true; + } + let text = self.current_chunk.text.as_ref(); + last_flush_time.elapsed() > FLUSH_INTERVAL || text.is_some_and(|t| t.len() > MAX_CHUNK_SIZE) + } + + /// Flushes the current chunk(s) to the Redis stream. + pub(super) async fn flush_chunks( + &mut self, + ws_writer: &mut WsWriter, + ) -> Result<(), StreamingError> { + let chunk_state = std::mem::take(&mut self.current_chunk); + + if let Some(text) = chunk_state.text { + self.add_to_stream(ws_writer, RedisStreamChunk::Text(text)) + .await?; + } + // if let Some(tool_calls) = chunk_state.tool_calls { + // for tool_call in tool_calls { + // let tool_call_str = serde_json::to_string(&tool_call).unwrap_or_default(); + // let entry = RedisStreamChunk::ToolCall(tool_call_str); + // self.add_to_stream(ws_writer, entry).await?; + // } + // } + // if let Some(pending_tool_calls) = chunk_state.pending_tool_calls { + // for tool_call in pending_tool_calls { + // let tool_call_str = serde_json::to_string(&tool_call).unwrap_or_default(); + // let entry = RedisStreamChunk::PendingToolCall(tool_call_str); + // self.add_to_stream(ws_writer, entry).await?; + // } + // } + if let Some(error) = chunk_state.error { + self.add_to_stream(ws_writer, RedisStreamChunk::Error(error)) + .await?; + } + + Ok(ws_writer.flush().await?) + } + + /// Serialize and add an entry to Redis via the WebSocket connection (does not flush the connection) + async fn add_to_stream( + &mut self, + ws_writer: &mut WsWriter, + entry: RedisStreamChunk, + ) -> Result<(), StreamingError> { + let message = WsMessage::text_from_json(&entry)?; + Ok(ws_writer.feed(message).await?) + } +} diff --git a/server-new/src/state.rs b/server-new/src/state.rs index bd34ce4..2ea3b92 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -10,6 +10,7 @@ use crate::{ services::{ auth::{AuthService, oauth::OAuthProviderMap}, chat::ChatService, + stream::tinistream::TinistreamClient, }, }; @@ -23,6 +24,7 @@ pub struct AppStateInner { pub http_client: reqwest::Client, pub db_pool: DbPool, pub redis: fred::prelude::Pool, + pub tinistream: TinistreamClient, pub oauth_providers: OAuthProviderMap, } @@ -37,7 +39,7 @@ impl AppState { } pub fn chat_service(&self) -> ChatService<'_> { - ChatService::new(&self.http_client, self.redis.next()) + ChatService::new(&self.db_pool, &self.tinistream) } } From 900b57587b928a0fefc5133303cce6af7f6f19f6 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 30 Jun 2026 02:12:04 -0400 Subject: [PATCH 046/111] chat service - http clients --- docker-compose.yml | 1 + server-new/config.toml | 4 +++ server-new/src/config.rs | 7 +++++ server-new/src/lib.rs | 5 +--- server-new/src/plugins/auth.rs | 6 ++-- server-new/src/plugins/clients.rs | 42 +++++++++++++++++++++++++++ server-new/src/plugins/mod.rs | 1 + server-new/src/services/auth/oauth.rs | 1 + 8 files changed, 60 insertions(+), 7 deletions(-) create mode 100644 server-new/src/plugins/clients.rs diff --git a/docker-compose.yml b/docker-compose.yml index 47e605d..687cc5e 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -27,6 +27,7 @@ services: - "8081:8081" environment: STREAMER_PORT: 8081 + STREAMER_API_KEY: dev-streamer-api-key STREAMER_SERVER_ADDRESS: http://localhost:8081 STREAMER_REDIS_URL: redis://redis:6379 STREAMER_TTL: 360 diff --git a/server-new/config.toml b/server-new/config.toml index 977cd6d..31d87ef 100644 --- a/server-new/config.toml +++ b/server-new/config.toml @@ -11,6 +11,10 @@ url = "postgres://localhost" [redis] url = "redis://localhost:6379" +[services] +streamer_url = "http://localhost:8081" +streamer_api_key = "dev-streamer-api-key" + [auth] cookie_name = "auth-rs-chat" session_length = 604800 diff --git a/server-new/src/config.rs b/server-new/src/config.rs index 76af488..b56d546 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -16,6 +16,7 @@ pub struct AppConfig { pub server: ServerConfig, pub database: DatabaseConfig, pub auth: AuthConfig, + pub services: ServiceConfig, pub security: SecurityConfig, pub redis: RedisConfig, } @@ -35,6 +36,12 @@ pub struct DatabaseConfig { pub url: String, } +#[derive(Debug, Clone, Deserialize)] +pub struct ServiceConfig { + pub streamer_url: String, + pub streamer_api_key: String, +} + #[derive(Debug, Clone, Deserialize)] pub struct AuthConfig { pub cookie_key: String, diff --git a/server-new/src/lib.rs b/server-new/src/lib.rs index fecbebc..0df07ba 100644 --- a/server-new/src/lib.rs +++ b/server-new/src/lib.rs @@ -13,12 +13,9 @@ mod services; mod state; pub async fn create_app() -> anyhow::Result> { - let http_client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build()?; let app = App::new() - .store(http_client) // Add shared http client .register(config::plugin()) // Extract configuration and add to state + .register(plugins::clients::plugin()) // Initialize HTTP clients .register(plugins::database::plugin()) // Initialize database .register(plugins::redis::plugin()) // Initialize Redis .register(api::plugin()) // Add API routes diff --git a/server-new/src/plugins/auth.rs b/server-new/src/plugins/auth.rs index 8d0f6cd..773fe9c 100644 --- a/server-new/src/plugins/auth.rs +++ b/server-new/src/plugins/auth.rs @@ -19,16 +19,16 @@ const CLEANUP_INTERVAL: Duration = Duration::from_mins(15); /// Add auth & session handling to the server. Sessions are stored in Postgres and cached in Redis. pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("Session") + AdHocPlugin::named("Auth") .on_init(async |mut state| { let config = state.get::().context("no config")?; - let db_pool = state.get::().context("no db pool")?.clone(); + let db_pool = state.get::().context("no db pool")?.to_owned(); // Build configured OAuth providers state.insert(OAuthService::build_provider_map(&config.auth)); // Session cleanup task - tokio::task::spawn(async move { + tokio::spawn(async move { let mut interval = tokio::time::interval(CLEANUP_INTERVAL); interval.tick().await; diff --git a/server-new/src/plugins/clients.rs b/server-new/src/plugins/clients.rs new file mode 100644 index 0000000..a3bb2d3 --- /dev/null +++ b/server-new/src/plugins/clients.rs @@ -0,0 +1,42 @@ +use std::time::Duration; + +use anyhow::Context; +use axum_plugin::AdHocPlugin; + +use crate::{config::AppConfig, services::stream::tinistream::TinistreamClient, state::AppState}; + +// Default timeout for HTTP requests +const TIMEOUT: Duration = Duration::from_secs(10); + +/// Setup HTTP clients for interacting with LLMs and services +pub fn plugin() -> AdHocPlugin { + AdHocPlugin::named("Clients").on_init(async |mut state| { + let config = state.get::().context("no config")?; + + // Main HTTP client for LLM provider and OAuth requests. + // No total request timeout to allow for long-lived streaming responses. + let http_client = reqwest::ClientBuilder::new() + .connect_timeout(TIMEOUT) + .redirect(reqwest::redirect::Policy::none()) + .build()?; + + // `tinistream` client with API key header + let tini_http_client = reqwest::ClientBuilder::new() + .connect_timeout(TIMEOUT) + .timeout(TIMEOUT) + .default_headers({ + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert("X-API-KEY", config.services.streamer_api_key.parse()?); + headers + }) + .build()?; + let tinistream = TinistreamClient::new(tinistream_client::Client::new_with_client( + &config.services.streamer_url, + tini_http_client, + )); + + state.insert(http_client); + state.insert(tinistream); + Ok(state) + }) +} diff --git a/server-new/src/plugins/mod.rs b/server-new/src/plugins/mod.rs index de05e46..278b9f2 100644 --- a/server-new/src/plugins/mod.rs +++ b/server-new/src/plugins/mod.rs @@ -1,4 +1,5 @@ pub mod auth; +pub mod clients; pub mod database; pub mod logging; pub mod redis; diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index 912955b..fd2fcbe 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -221,6 +221,7 @@ impl<'a> OAuthService<'a> { discord::DiscordProvider, github::GitHubProvider, google::GoogleProvider, oidc::OidcProvider, }; + let mut map: OAuthProviderMap = HashMap::new(); if let Some(ref c) = config.github { map.insert(OAuthProviderEnum::Github, Box::new(GitHubProvider::new(c))); From dd74bb0ce8b839ea53df14a13cd8bcb75d73564b Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 1 Jul 2026 04:40:07 -0400 Subject: [PATCH 047/111] chat service - full flow --- server-new/.env.example | 2 +- server-new/Cargo.lock | 201 ++++++++++++++++-- server-new/Cargo.toml | 1 + server-new/src/api/auth.rs | 17 +- server-new/src/api/chat.rs | 58 +++++ server-new/src/api/hello.rs | 15 -- server-new/src/api/mod.rs | 4 +- server-new/src/config.rs | 4 +- server-new/src/db/mod.rs | 7 +- server-new/src/db/models.rs | 8 +- server-new/src/db/models/provider.rs | 95 +++++++++ server-new/src/db/models/secret.rs | 46 ++++ server-new/src/db/repositories.rs | 2 + server-new/src/db/repositories/chat.rs | 21 +- server-new/src/db/repositories/provider.rs | 98 +++++++++ server-new/src/extractors/database.rs | 20 ++ server-new/src/extractors/mod.rs | 1 + server-new/src/extractors/session.rs | 39 +++- server-new/src/llm/error.rs | 4 +- server-new/src/llm/interface.rs | 7 +- server-new/src/llm/providers/openai/mod.rs | 146 ++++++------- server-new/src/llm/types.rs | 8 +- server-new/src/plugins/auth.rs | 28 ++- server-new/src/services/auth/encryption.rs | 76 +++++++ server-new/src/services/auth/error.rs | 4 +- server-new/src/services/auth/mod.rs | 15 +- server-new/src/services/auth/oauth.rs | 12 +- server-new/src/services/auth/session.rs | 8 +- server-new/src/services/auth/session_store.rs | 8 +- server-new/src/services/auth/types.rs | 26 --- server-new/src/services/chat/error.rs | 22 +- server-new/src/services/chat/mod.rs | 70 ++++-- server-new/src/services/chat/titles.rs | 74 +++++++ server-new/src/services/mod.rs | 1 + server-new/src/services/provider/error.rs | 27 +++ server-new/src/services/provider/mod.rs | 80 +++++++ server-new/src/services/stream/error.rs | 2 +- server-new/src/services/stream/mod.rs | 5 + server-new/src/services/stream/writer.rs | 4 +- server-new/src/state.rs | 19 +- 40 files changed, 1024 insertions(+), 261 deletions(-) create mode 100644 server-new/src/api/chat.rs delete mode 100644 server-new/src/api/hello.rs create mode 100644 server-new/src/db/models/provider.rs create mode 100644 server-new/src/db/models/secret.rs create mode 100644 server-new/src/db/repositories/provider.rs create mode 100644 server-new/src/extractors/database.rs create mode 100644 server-new/src/services/auth/encryption.rs delete mode 100644 server-new/src/services/auth/types.rs create mode 100644 server-new/src/services/chat/titles.rs create mode 100644 server-new/src/services/provider/error.rs create mode 100644 server-new/src/services/provider/mod.rs diff --git a/server-new/.env.example b/server-new/.env.example index 9660a89..a2962b6 100644 --- a/server-new/.env.example +++ b/server-new/.env.example @@ -3,7 +3,7 @@ RS_CHAT_DATABASE__URL=postgres://postgres:postgres@localhost/postgres DATABASE_URL=postgres://postgres:postgres@localhost/postgres # Auth -RS_CHAT_AUTH__COOKIE_KEY= # hex secret >=32 bytes, e.g. `openssl rand --hex 32` +RS_CHAT_AUTH__ENCRYPTION_KEY= # 32-byte hex secret, e.g. `openssl rand --hex 32` RS_CHAT_AUTH__GITHUB__CLIENT_ID= RS_CHAT_AUTH__GITHUB__CLIENT_SECRET= diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index 0a5d9ab..dcecdfe 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -8,10 +8,20 @@ version = "0.5.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0" dependencies = [ - "crypto-common", + "crypto-common 0.1.7", "generic-array", ] +[[package]] +name = "aead" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1973cfbc1a2daf9cf550e74e1f088c28e7f7d8c1e1418fb6c9dc5184b7e84c99" +dependencies = [ + "crypto-common 0.2.2", + "inout 0.2.2", +] + [[package]] name = "aes" version = "0.8.4" @@ -19,8 +29,19 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0" dependencies = [ "cfg-if", - "cipher", - "cpufeatures", + "cipher 0.4.4", + "cpufeatures 0.2.17", +] + +[[package]] +name = "aes" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1fc76eaeac4c9164506c466d4ffdd8ec9d0c5bf57ee97177c4d8eceb3a0e138" +dependencies = [ + "cipher 0.5.2", + "cpubits", + "cpufeatures 0.3.0", ] [[package]] @@ -29,11 +50,25 @@ version = "0.10.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "831010a0f742e1209b3bcea8fab6a8e149051ba6099432c8cb2cc117dec3ead1" dependencies = [ - "aead", - "aes", - "cipher", - "ctr", - "ghash", + "aead 0.5.2", + "aes 0.8.4", + "cipher 0.4.4", + "ctr 0.9.2", + "ghash 0.5.1", + "subtle", +] + +[[package]] +name = "aes-gcm" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fdf011db2e21ce0d575593d749db5554b47fed37aff429e4dc50bc91ac93a028" +dependencies = [ + "aead 0.6.1", + "aes 0.9.1", + "cipher 0.5.2", + "ctr 0.10.1", + "ghash 0.6.0", "subtle", ] @@ -270,6 +305,15 @@ dependencies = [ "generic-array", ] +[[package]] +name = "block-buffer" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2f6c7dbe95a6ed67ad9f18e57daf93a2f034c524b99fd2b76d18fdfeb6660aa" +dependencies = [ + "hybrid-array", +] + [[package]] name = "bon" version = "3.9.3" @@ -373,8 +417,19 @@ version = "0.4.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" dependencies = [ - "crypto-common", - "inout", + "crypto-common 0.1.7", + "inout 0.1.4", +] + +[[package]] +name = "cipher" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e8cf2a2c93cd704877c0858356ed03480ff301ee950b43f1cbe4573b088bfa6c" +dependencies = [ + "block-buffer 0.12.1", + "crypto-common 0.2.2", + "inout 0.2.2", ] [[package]] @@ -386,6 +441,12 @@ dependencies = [ "cc", ] +[[package]] +name = "cmov" +version = "0.5.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c9ea0ac24bc397ab3c98583a3c9ba74fa56b09a4449bbe172b9b1ddb016027a" + [[package]] name = "combine" version = "4.6.7" @@ -402,7 +463,7 @@ version = "0.18.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4ddef33a339a91ea89fb53151bd0a4689cfce27055c291dfa69945475d22c747" dependencies = [ - "aes-gcm", + "aes-gcm 0.10.3", "base64", "hkdf", "hmac", @@ -436,6 +497,12 @@ version = "0.8.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "773648b94d0e5d620f64f280777445740e61fe701025087ec8b57f45c791888b" +[[package]] +name = "cpubits" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "15b85f9c39137c3a891689859392b1bd49812121d0d61c9caf00d46ed5ce06ae" + [[package]] name = "cpufeatures" version = "0.2.17" @@ -445,6 +512,15 @@ dependencies = [ "libc", ] +[[package]] +name = "cpufeatures" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" +dependencies = [ + "libc", +] + [[package]] name = "crc16" version = "0.4.0" @@ -477,13 +553,42 @@ dependencies = [ "typenum", ] +[[package]] +name = "crypto-common" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ce6e4c961d6cd6c9a86db418387425e8bdeaf05b3c8bc1411e6dca4c252f1453" +dependencies = [ + "getrandom 0.4.3", + "hybrid-array", + "rand_core 0.10.1", +] + [[package]] name = "ctr" version = "0.9.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0369ee1ad671834580515889b80f2ea915f23b8be8d0daa4bbaf2ac5c7590835" dependencies = [ - "cipher", + "cipher 0.4.4", +] + +[[package]] +name = "ctr" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "baaca1c4b237092596f64d571e9db6ce4109c4ef9742e27590f1709594461f21" +dependencies = [ + "cipher 0.5.2", +] + +[[package]] +name = "ctutils" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7d5515a3834141de9eafb9717ad39eea8247b5674e6066c404e8c4b365d2a29e" +dependencies = [ + "cmov", ] [[package]] @@ -679,8 +784,8 @@ version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ - "block-buffer", - "crypto-common", + "block-buffer 0.10.4", + "crypto-common 0.1.7", "subtle", ] @@ -976,6 +1081,7 @@ dependencies = [ "cfg-if", "libc", "r-efi 6.0.0", + "rand_core 0.10.1", ] [[package]] @@ -985,7 +1091,16 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f0d8a4362ccb29cb0b265253fb0a2728f592895ee6854fd9bc13f2ffda266ff1" dependencies = [ "opaque-debug", - "polyval", + "polyval 0.6.2", +] + +[[package]] +name = "ghash" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2eecf2d5dc9b66b732b97707a0210906b1d30523eb773193ab777c0c84b3e8d5" +dependencies = [ + "polyval 0.7.1", ] [[package]] @@ -1087,6 +1202,15 @@ version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" +[[package]] +name = "hybrid-array" +version = "0.4.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "818356c5132c1fede50f837ca96afbe78ff42413047f4abb886217845e1b6c8c" +dependencies = [ + "typenum", +] + [[package]] name = "hyper" version = "1.10.1" @@ -1304,6 +1428,15 @@ dependencies = [ "generic-array", ] +[[package]] +name = "inout" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4250ce6452e92010fdf7268ccc5d14faa80bb12fc741938534c58f16804e03c7" +dependencies = [ + "hybrid-array", +] + [[package]] name = "ipnet" version = "2.12.0" @@ -1702,9 +1835,20 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.2.17", "opaque-debug", - "universal-hash", + "universal-hash 0.5.1", +] + +[[package]] +name = "polyval" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7dfc63250416fea14f5749b90725916a6c903f599d51cb635aa7a52bfd03eede" +dependencies = [ + "cpubits", + "cpufeatures 0.3.0", + "universal-hash 0.6.1", ] [[package]] @@ -1943,6 +2087,12 @@ dependencies = [ "getrandom 0.3.4", ] +[[package]] +name = "rand_core" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" + [[package]] name = "redis-protocol" version = "6.0.0" @@ -2081,6 +2231,7 @@ dependencies = [ name = "rs-chat-api" version = "0.1.0" dependencies = [ + "aes-gcm 0.11.0", "anyhow", "async-stream", "async-trait", @@ -2387,7 +2538,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.2.17", "digest", ] @@ -2398,7 +2549,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.2.17", "digest", ] @@ -3153,10 +3304,20 @@ version = "0.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea" dependencies = [ - "crypto-common", + "crypto-common 0.1.7", "subtle", ] +[[package]] +name = "universal-hash" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f4987bdc12753382e0bec4a65c50738ffaabc998b9cdd1f952fb5f39b0048a96" +dependencies = [ + "crypto-common 0.2.2", + "ctutils", +] + [[package]] name = "untrusted" version = "0.9.0" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 3a540f1..ea6819d 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -6,6 +6,7 @@ description = "RsChat Server" publish = false [dependencies] +aes-gcm = "0.11.0" anyhow = "1.0.102" async-stream = "0.3.6" async-trait = "0.1.89" diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index c881f8d..13f94fa 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -10,7 +10,11 @@ use serde::Deserialize; use crate::{ api::RoutePrefix, error::AppResult, - services::auth::{SessionMeta, UserSession, oauth::OAuthProviderEnum}, + extractors::{ + database::Database, + session::{SessionMeta, UserSession}, + }, + services::auth::oauth::OAuthProviderEnum, state::AppState, }; @@ -50,6 +54,7 @@ async fn callback_handler( Path(provider): Path, Query(query): Query, Extension(RoutePrefix(prefix)): Extension, + Database(mut db): Database, State(state): State, session: tower_sessions::Session, meta: SessionMeta, @@ -65,22 +70,24 @@ async fn callback_handler( &query.state, ) .await?; - let user = oauth.get_user(provider, &token, maybe_user).await?; + let user = oauth + .get_user(&mut db, provider, &token, maybe_user) + .await?; state .auth_service() .session() .login(&session, &meta, &user.id) .await?; - Ok(Redirect::to("/api/auth/user")) - // Ok(Redirect::to(&state.config.server.base_url)) + Ok(Redirect::to(&state.config.server.base_url)) } async fn get_user_handler( UserSession { user_id }: UserSession, + Database(mut db): Database, State(state): State, ) -> AppResult { - let user = state.auth_service().get_user(&user_id).await?; + let user = state.auth_service().get_user(&mut db, &user_id).await?; Ok(Json(user)) } diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs new file mode 100644 index 0000000..b8816c6 --- /dev/null +++ b/server-new/src/api/chat.rs @@ -0,0 +1,58 @@ +use axum::{ + Json, + extract::{Path, State}, + response::IntoResponse, + routing, +}; +use serde::Deserialize; +use uuid::Uuid; + +use crate::{ + error::AppError, + extractors::{database::Database, session::UserSession}, + llm::types::{LlmChatOptions, LlmUserMessage}, + state::AppState, +}; + +pub fn routes() -> axum::Router { + axum::Router::new().route("/{session_id}", routing::post(chat_stream)) +} + +#[derive(Debug, Deserialize)] +struct ChatInput { + /// The new chat message from the user + message: Option, + /// The ID of the provider to chat with + provider_id: i32, + /// Configuration for the provider + options: LlmChatOptions, +} + +async fn chat_stream( + UserSession { user_id }: UserSession, + Path(session_id): Path, + Database(mut db): Database, + State(state): State, + Json(input): Json, +) -> Result { + let llm_provider = state + .provider_service() + .build_llm_provider(&mut db, &user_id, input.provider_id) + .await?; + let stream_access = state + .chat_service() + .stream_user_chat( + &mut db, + user_id, + session_id, + input.provider_id, + llm_provider, + input + .message + .map(|text| LlmUserMessage { text, files: None }), + input.options, + ) + .await?; + + Ok(Json(stream_access)) +} diff --git a/server-new/src/api/hello.rs b/server-new/src/api/hello.rs deleted file mode 100644 index 52fd573..0000000 --- a/server-new/src/api/hello.rs +++ /dev/null @@ -1,15 +0,0 @@ -use crate::{error::AppResult, state::AppState}; - -pub fn routes() -> axum::Router { - axum::Router::new() - .route("/", axum::routing::get(hello_handler)) - .route("/", axum::routing::post(post_handler)) -} - -async fn hello_handler() -> AppResult { - Ok("Hello, World!".to_string()) -} - -async fn post_handler() -> AppResult { - Ok("Post handler!".to_string()) -} diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index 38f55d1..a11a1f6 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -4,8 +4,8 @@ use axum_plugin::AdHocPlugin; use crate::state::AppState; pub mod auth; +pub mod chat; pub mod health; -pub mod hello; /// Adds all API routes to the server under `/api` pub fn plugin() -> AdHocPlugin { @@ -15,7 +15,7 @@ pub fn plugin() -> AdHocPlugin { "/auth", auth::routes().layer(Extension(RoutePrefix("/api/auth"))), ) - .nest("/hello", hello::routes()) + .nest("/chat", chat::routes()) .nest("/health", health::routes()); Ok(router.nest("/api", api_routes)) diff --git a/server-new/src/config.rs b/server-new/src/config.rs index b56d546..7b104dc 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -44,7 +44,7 @@ pub struct ServiceConfig { #[derive(Debug, Clone, Deserialize)] pub struct AuthConfig { - pub cookie_key: String, + pub encryption_key: String, pub cookie_name: String, pub session_length: i64, pub github: Option, @@ -78,7 +78,7 @@ pub fn plugin() -> AdHocPlugin { }) } -/// Extract configuration from config.toml, then environment overrides. +/// Extract configuration from config.toml, then environment variables. fn extract_config() -> anyhow::Result { let config = figment::Figment::new() .merge(Toml::file("config.toml")) diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs index d597ace..cab508e 100644 --- a/server-new/src/db/mod.rs +++ b/server-new/src/db/mod.rs @@ -21,7 +21,7 @@ pub struct DbConnection(Object); impl Deref for DbConnection { type Target = AsyncPgConnection; fn deref(&self) -> &Self::Target { - self.0.as_ref() + &self.0 } } impl DerefMut for DbConnection { @@ -51,10 +51,13 @@ impl DbService { pub fn users(&mut self) -> repositories::UserRepository<'_> { repositories::UserRepository::new(&mut self.cxn) } - pub fn sessions(&mut self) -> repositories::SessionRepository<'_> { + pub fn auth_sessions(&mut self) -> repositories::SessionRepository<'_> { repositories::SessionRepository::new(&mut self.cxn) } pub fn chats(&mut self) -> repositories::ChatRepository<'_> { repositories::ChatRepository::new(&mut self.cxn) } + pub fn providers(&mut self) -> repositories::ProviderRepository<'_> { + repositories::ProviderRepository::new(&mut self.cxn) + } } diff --git a/server-new/src/db/models.rs b/server-new/src/db/models.rs index 80c3d61..feae9a2 100644 --- a/server-new/src/db/models.rs +++ b/server-new/src/db/models.rs @@ -3,8 +3,8 @@ use crate::db::schema; // mod api_key; mod chat; // mod file; -// mod provider; -// mod secret; +mod provider; +mod secret; // mod tool; mod session; mod user; @@ -12,8 +12,8 @@ mod user; // pub use api_key::*; pub use chat::*; // pub use file::*; -// pub use provider::*; -// pub use secret::*; +pub use provider::*; +pub use secret::*; // pub use tool::*; pub use session::*; pub use user::*; diff --git a/server-new/src/db/models/provider.rs b/server-new/src/db/models/provider.rs new file mode 100644 index 0000000..49411fe --- /dev/null +++ b/server-new/src/db/models/provider.rs @@ -0,0 +1,95 @@ +use std::str::FromStr; + +use chrono::{DateTime, Utc}; +use diesel::prelude::*; +use serde::{Deserialize, Serialize}; +use uuid::Uuid; + +use crate::db::models::ChatRsUser; + +#[derive(Identifiable, Associations, Queryable, Selectable, Serialize)] +#[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] +#[diesel(table_name = super::schema::providers)] +pub struct ChatRsProvider { + pub id: i32, + pub name: String, + // #[schemars(with = "ChatRsProviderType")] + pub provider_type: String, + // #[schemars(with = "OpenaiSubtype")] + // pub openai_subtype: Option, + #[serde(skip)] + pub user_id: Uuid, + pub default_model: String, + pub base_url: Option, + pub api_key_id: Option, + pub created_at: DateTime, +} + +#[derive(Insertable)] +#[diesel(table_name = super::schema::providers)] +pub struct NewChatRsProvider<'a> { + pub name: &'a str, + pub provider_type: &'a str, + pub user_id: &'a Uuid, + pub base_url: Option<&'a str>, + pub default_model: &'a str, + pub api_key_id: Option, +} + +#[derive(Default, AsChangeset)] +#[diesel(table_name = super::schema::providers)] +pub struct UpdateChatRsProvider<'a> { + pub name: Option<&'a str>, + pub base_url: Option<&'a str>, + pub default_model: Option<&'a str>, + pub api_key_id: Option, +} + +/// The API type of the provider +#[derive(Debug, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ChatRsProviderType { + Anthropic, + Openai, + Ollama, + Lorem, +} + +/// The subtype for OpenAI-compatible providers +#[derive(Debug, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum OpenaiSubtype { + Openai, + Google, + OpenRouter, + LlmGateway, +} + +#[derive(Debug, thiserror::Error)] +#[error("invalid provider type: '{0}'")] +pub struct ParseProviderTypeError(String); + +impl FromStr for ChatRsProviderType { + type Err = ParseProviderTypeError; + + fn from_str(value: &str) -> Result { + match value { + "anthropic" => Ok(ChatRsProviderType::Anthropic), + "openai" => Ok(ChatRsProviderType::Openai), + "ollama" => Ok(ChatRsProviderType::Ollama), + "lorem" => Ok(ChatRsProviderType::Lorem), + provider => Err(ParseProviderTypeError(provider.into())), + } + } +} + +impl From<&ChatRsProviderType> for &str { + fn from(value: &ChatRsProviderType) -> Self { + match value { + ChatRsProviderType::Anthropic => "anthropic", + ChatRsProviderType::Openai => "openai", + ChatRsProviderType::Ollama => "ollama", + ChatRsProviderType::Lorem => "lorem", + } + } +} diff --git a/server-new/src/db/models/secret.rs b/server-new/src/db/models/secret.rs new file mode 100644 index 0000000..dec5164 --- /dev/null +++ b/server-new/src/db/models/secret.rs @@ -0,0 +1,46 @@ +use chrono::{DateTime, Utc}; +use diesel::prelude::*; +use serde::Serialize; +use uuid::Uuid; + +use crate::db::models::ChatRsUser; + +#[derive(Identifiable, Queryable, Selectable, Associations)] +#[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] +#[diesel(table_name = super::schema::secrets)] +pub struct ChatRsSecret { + pub id: Uuid, + pub user_id: Uuid, + pub name: String, + pub ciphertext: Vec, + pub nonce: Vec, + pub created_at: DateTime, +} + +#[derive(Identifiable, Queryable, Selectable, Associations, Serialize)] +#[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] +#[diesel(table_name = super::schema::secrets)] +pub struct ChatRsSecretMeta { + pub id: Uuid, + #[serde(skip)] + pub user_id: Uuid, + pub name: String, + pub created_at: DateTime, +} + +#[derive(Insertable)] +#[diesel(table_name = super::schema::secrets)] +pub struct NewChatRsSecret<'r> { + pub user_id: &'r Uuid, + pub name: &'r str, + pub ciphertext: &'r Vec, + pub nonce: &'r Vec, +} + +#[derive(Default, AsChangeset)] +#[diesel(table_name = super::schema::secrets)] +pub struct UpdateChatRsSecret<'r> { + pub name: Option<&'r str>, + pub ciphertext: Option<&'r Vec>, + pub nonce: Option<&'r Vec>, +} diff --git a/server-new/src/db/repositories.rs b/server-new/src/db/repositories.rs index 959d916..16624ea 100644 --- a/server-new/src/db/repositories.rs +++ b/server-new/src/db/repositories.rs @@ -1,7 +1,9 @@ mod chat; +mod provider; mod session; mod user; pub use chat::ChatRepository; +pub use provider::ProviderRepository; pub use session::SessionRepository; pub use user::UserRepository; diff --git a/server-new/src/db/repositories/chat.rs b/server-new/src/db/repositories/chat.rs index 9de8566..d7237af 100644 --- a/server-new/src/db/repositories/chat.rs +++ b/server-new/src/db/repositories/chat.rs @@ -1,5 +1,3 @@ -use std::ops::{Deref, DerefMut}; - use diesel::prelude::*; use diesel_async::RunQueryDsl; use uuid::Uuid; @@ -86,7 +84,7 @@ impl<'a> ChatRepository<'a> { Ok(id.to_string()) } - pub async fn get_all_sessions( + pub async fn get_recent_sessions( &mut self, user_id: &Uuid, ) -> Result, diesel::result::Error> { @@ -101,27 +99,28 @@ impl<'a> ChatRepository<'a> { Ok(sessions) } - pub async fn get_session( + pub async fn find_session( &mut self, user_id: &Uuid, session_id: &Uuid, - ) -> Result { + ) -> Result, diesel::result::Error> { let session = chat_sessions::table .filter(chat_sessions::user_id.eq(user_id)) .filter(chat_sessions::id.eq(session_id)) .select(ChatRsSession::as_select()) .first(self.db) - .await?; + .await + .optional()?; Ok(session) } - pub async fn get_session_with_messages( + pub async fn find_session_with_messages( &mut self, user_id: &Uuid, session_id: &Uuid, - ) -> Result<(ChatRsSession, Vec), diesel::result::Error> { - let (session, messages) = futures::future::try_join( + ) -> Result<(Option, Vec), diesel::result::Error> { + let (session, messages) = futures::future::join( chat_sessions::table .filter(chat_sessions::user_id.eq(user_id)) .filter(chat_sessions::id.eq(session_id)) @@ -133,9 +132,9 @@ impl<'a> ChatRepository<'a> { .order_by(chat_messages::created_at.asc()) .load(&mut &**self.db), ) - .await?; + .await; - Ok((session, messages)) + Ok((session.optional()?, messages?)) } pub async fn search_sessions( diff --git a/server-new/src/db/repositories/provider.rs b/server-new/src/db/repositories/provider.rs new file mode 100644 index 0000000..0459a63 --- /dev/null +++ b/server-new/src/db/repositories/provider.rs @@ -0,0 +1,98 @@ +use diesel::prelude::*; +use diesel_async::RunQueryDsl; +use uuid::Uuid; + +use crate::db::{ + DbConnection, + models::{ChatRsProvider, ChatRsSecret, NewChatRsProvider, UpdateChatRsProvider}, + schema::{providers, secrets}, +}; + +pub struct ProviderRepository<'a> { + pub db: &'a mut DbConnection, +} + +impl<'a> ProviderRepository<'a> { + pub fn new(db: &'a mut DbConnection) -> Self { + ProviderRepository { db } + } + + pub async fn find_by_id( + &mut self, + user_id: &Uuid, + provider_id: i32, + ) -> Result)>, diesel::result::Error> { + providers::table + .left_join(secrets::table) + .filter(providers::user_id.eq(user_id)) + .filter(providers::id.eq(provider_id)) + .select(( + ChatRsProvider::as_select(), + Option::::as_select(), + )) + .first(self.db) + .await + .optional() + } + + pub async fn list_by_user_id( + &mut self, + user_id: &Uuid, + ) -> Result, diesel::result::Error> { + providers::table + .filter(providers::user_id.eq(user_id)) + .select(ChatRsProvider::as_select()) + .load(self.db) + .await + } + + pub async fn create( + &mut self, + provider: NewChatRsProvider<'_>, + ) -> Result { + diesel::insert_into(providers::table) + .values(provider) + .returning(ChatRsProvider::as_returning()) + .get_result(self.db) + .await + } + + pub async fn update( + &mut self, + user_id: &Uuid, + provider_id: i32, + data: UpdateChatRsProvider<'_>, + ) -> Result { + diesel::update(providers::table) + .filter(providers::user_id.eq(user_id)) + .filter(providers::id.eq(provider_id)) + .set(data) + .returning(ChatRsProvider::as_returning()) + .get_result(self.db) + .await + } + + pub async fn delete( + &mut self, + user_id: &Uuid, + provider_id: i32, + ) -> Result { + diesel::delete(providers::table) + .filter(providers::user_id.eq(user_id)) + .filter(providers::id.eq(provider_id)) + .returning(ChatRsProvider::as_returning()) + .get_result(self.db) + .await + } + + pub async fn delete_by_user( + &mut self, + user_id: &Uuid, + ) -> Result, diesel::result::Error> { + diesel::delete(providers::table) + .filter(providers::user_id.eq(user_id)) + .returning(ChatRsProvider::as_returning()) + .get_results(self.db) + .await + } +} diff --git a/server-new/src/extractors/database.rs b/server-new/src/extractors/database.rs new file mode 100644 index 0000000..fccfda1 --- /dev/null +++ b/server-new/src/extractors/database.rs @@ -0,0 +1,20 @@ +use axum::extract::FromRequestParts; + +use crate::{db::DbService, error::AppError, state::AppState}; + +/// An extractor to retrieve a database connection from the pool +pub struct Database(pub DbService); + +impl FromRequestParts for Database { + type Rejection = AppError; + + async fn from_request_parts( + _parts: &mut axum::http::request::Parts, + state: &AppState, + ) -> Result { + match DbService::from_pool(&state.db_pool).await { + Ok(db_service) => Ok(Self(db_service)), + Err(err) => Err(AppError::internal(err.into())), + } + } +} diff --git a/server-new/src/extractors/mod.rs b/server-new/src/extractors/mod.rs index f52f1c4..26d219c 100644 --- a/server-new/src/extractors/mod.rs +++ b/server-new/src/extractors/mod.rs @@ -1 +1,2 @@ +pub mod database; pub mod session; diff --git a/server-new/src/extractors/session.rs b/server-new/src/extractors/session.rs index 2293cd8..c902367 100644 --- a/server-new/src/extractors/session.rs +++ b/server-new/src/extractors/session.rs @@ -9,19 +9,38 @@ use axum::{ }; use tower_sessions::Session; -use crate::{ - error::AppError, - services::auth::{SessionMeta, UserSession}, - state::AppState, -}; +use crate::{error::AppError, state::AppState}; -/// Active user session data. -/// -/// This can be used as an extractor in route handlers: -/// - If used as `Option`, will be `Some` if there is an active session -/// and `None` otherwise. +use serde::{Deserialize, Serialize}; +use uuid::Uuid; + +use crate::db::UtcDateTime; + +/// Represents an active user session. This can be used as an extractor in route handlers: /// - If used as `UserSession`, request will automatically return an unauthorized error /// if there is no active session. +/// - If used as `Option`, will be `Some` if there is an active session +/// and `None` otherwise. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct UserSession { + pub user_id: Uuid, +} + +impl UserSession { + pub fn new(user_id: Uuid) -> Self { + Self { user_id } + } +} + +/// Session metadata extracted on login. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SessionMeta { + pub start_time: UtcDateTime, + pub ip: Option, + pub user_agent: Option, +} + +/// Active user session data. impl OptionalFromRequestParts for UserSession { type Rejection = AppError; diff --git a/server-new/src/llm/error.rs b/server-new/src/llm/error.rs index 61575d7..b68d66b 100644 --- a/server-new/src/llm/error.rs +++ b/server-new/src/llm/error.rs @@ -3,8 +3,10 @@ use crate::services::stream::error::StreamingError; /// Errors that can occur in an LLM provider request #[derive(Debug, thiserror::Error)] pub enum LlmRequestError { - #[error("Provider error: {0}")] + #[error("provider error: {0}")] Provider(String), + #[error("no content")] + NoContent, } /// Errors that can occur in an LLM stream chunk diff --git a/server-new/src/llm/interface.rs b/server-new/src/llm/interface.rs index d053c60..f594d8d 100644 --- a/server-new/src/llm/interface.rs +++ b/server-new/src/llm/interface.rs @@ -1,15 +1,20 @@ use futures::{future::BoxFuture, stream::BoxStream}; +use crate::llm::types::LlmPrompt; + use super::{ error::{LlmRequestError, LlmStreamChunkError}, types::{LlmChatRequest, LlmUsage}, }; -/// Trait for all LLM providers +/// Trait representing an LLM provider pub trait LlmProvider: Send + Sync { + fn prompt<'r>(&'r self, prompt: LlmPrompt<'r>) -> LlmPromptResponse<'r>; fn stream_chat<'r>(&'r self, request: LlmChatRequest<'r>) -> LlmStreamingResponse<'r>; } +/// API response to a prompt request from the LLM provider +pub type LlmPromptResponse<'r> = BoxFuture<'r, Result>; /// Initial API response to a streaming request from the LLM provider pub type LlmStreamingResponse<'r> = BoxFuture<'r, Result>; /// The response stream from the LLM provider diff --git a/server-new/src/llm/providers/openai/mod.rs b/server-new/src/llm/providers/openai/mod.rs index 7c0da58..7831818 100644 --- a/server-new/src/llm/providers/openai/mod.rs +++ b/server-new/src/llm/providers/openai/mod.rs @@ -4,9 +4,9 @@ use futures::StreamExt; use crate::llm::{ error::LlmRequestError, - interface::{LlmProvider, LlmStreamingResponse}, + interface::{LlmPromptResponse, LlmProvider, LlmStreamingResponse}, providers::utils, - types::LlmChatRequest, + types::{LlmChatRequest, LlmPrompt, LlmUsage}, }; mod request; @@ -85,41 +85,23 @@ impl OpenAIProviderConfig { #[derive(Debug, Clone)] pub struct OpenAIProvider { client: reqwest::Client, - _redis: fred::clients::Client, config: OpenAIProviderConfig, } impl OpenAIProvider { - pub fn new( - http_client: &reqwest::Client, - redis: &fred::clients::Client, - config: OpenAIProviderConfig, - ) -> Self { + pub fn new(http_client: &reqwest::Client, config: OpenAIProviderConfig) -> Self { Self { client: http_client.clone(), - _redis: redis.clone(), config, } } - pub fn openai( - http_client: &reqwest::Client, - redis: &fred::clients::Client, - api_key: impl Into, - ) -> Self { - Self::new(http_client, redis, OpenAIProviderConfig::openai(api_key)) + pub fn openai(http_client: &reqwest::Client, api_key: impl Into) -> Self { + Self::new(http_client, OpenAIProviderConfig::openai(api_key)) } - pub fn openrouter( - http_client: &reqwest::Client, - redis: &fred::clients::Client, - api_key: impl Into, - ) -> Self { - Self::new( - http_client, - redis, - OpenAIProviderConfig::openrouter(api_key), - ) + pub fn openrouter(http_client: &reqwest::Client, api_key: impl Into) -> Self { + Self::new(http_client, OpenAIProviderConfig::openrouter(api_key)) } } @@ -160,6 +142,59 @@ impl OpenAIRequestPolicy { } impl LlmProvider for OpenAIProvider { + fn prompt<'r>(&'r self, prompt: LlmPrompt<'r>) -> LlmPromptResponse<'r> { + let policy = OpenAIRequestPolicy::new(self.config.flavor); + let request = OpenAIRequest { + model: &prompt.options.model, + messages: vec![OpenAIMessage { + role: "user", + content: Some(vec![OpenAIContent::Text { text: prompt.text }]), + ..Default::default() + }], + max_tokens: policy.max_tokens(prompt.options.max_tokens), + max_completion_tokens: policy.max_completion_tokens(prompt.options.max_tokens), + store: policy.store(), + ..Default::default() + }; + + Box::pin(async move { + let provider_name = self.config.flavor.name(); + let response = self + .client + .post(format!("{}/chat/completions", self.config.base_url)) + .bearer_auth(&self.config.api_key) + .json(&request) + .send() + .await + .map_err(|err| { + LlmRequestError::Provider(format!("{provider_name} request failed: {err}")) + })?; + if !response.status().is_success() { + let status = response.status(); + let error_text = response.text().await.unwrap_or_default(); + return Err(LlmRequestError::Provider(format!( + "{provider_name} API error {status}: {error_text}", + ))); + } + + let mut openai_response: OpenAIResponse = response.json().await.map_err(|err| { + LlmRequestError::Provider(format!("Failed to parse response: {err}")) + })?; + let text = openai_response + .choices + .get_mut(0) + .and_then(|choice| choice.message.as_mut()) + .and_then(|message| message.content.take()) + .ok_or(LlmRequestError::NoContent)?; + if let Some(usage) = openai_response.usage { + let usage: LlmUsage = usage.into(); + tracing::info!("Prompt usage: {usage:?}"); + } + + Ok(text) + }) + } + fn stream_chat<'r>(&'r self, req: LlmChatRequest<'r>) -> LlmStreamingResponse<'r> { let policy = OpenAIRequestPolicy::new(self.config.flavor); let openai_messages = build_openai_messages(&req.messages); @@ -184,15 +219,13 @@ impl LlmProvider for OpenAIProvider { let response = self .client .post(format!("{}/chat/completions", self.config.base_url)) - .header("authorization", format!("Bearer {}", self.config.api_key)) - .header("content-type", "application/json") + .bearer_auth(&self.config.api_key) .json(&request) .send() .await .map_err(|e| { LlmRequestError::Provider(format!("{provider_name} request failed: {e}")) })?; - if !response.status().is_success() { let status = response.status(); let error_text = response.text().await.unwrap_or_default(); @@ -229,63 +262,6 @@ impl LlmProvider for OpenAIProvider { }) } - // async fn prompt( - // &self, - // message: &str, - // options: &LlmProviderOptions, - // ) -> Result { - // let request = OpenAIRequest { - // model: &options.model, - // messages: vec![OpenAIMessage { - // role: "user", - // content: Some(vec![OpenAIContent::Text { text: message }]), - // ..Default::default() - // }], - // max_tokens: options.max_tokens, - // temperature: options.temperature, - // store: (self.base_url == OPENAI_API_BASE_URL).then_some(false), - // ..Default::default() - // }; - - // let response = self - // .client - // .post(format!("{}/chat/completions", self.base_url)) - // .header("authorization", format!("Bearer {}", self.api_key)) - // .header("content-type", "application/json") - // .json(&request) - // .send() - // .await - // .map_err(|e| LlmError::ProviderError(format!("OpenAI request failed: {}", e)))?; - - // if !response.status().is_success() { - // let status = response.status(); - // let error_text = response.text().await.unwrap_or_default(); - // return Err(LlmError::ProviderError(format!( - // "OpenAI API error {}: {}", - // status, error_text - // ))); - // } - - // let mut openai_response: OpenAIResponse = response - // .json() - // .await - // .map_err(|e| LlmError::ProviderError(format!("Failed to parse response: {}", e)))?; - - // let text = openai_response - // .choices - // .get_mut(0) - // .and_then(|choice| choice.message.as_mut()) - // .and_then(|message| message.content.take()) - // .ok_or(LlmError::NoResponse)?; - - // if let Some(usage) = openai_response.usage { - // let usage: LlmUsage = usage.into(); - // println!("Prompt usage: {:?}", usage); - // } - - // Ok(text) - // } - // async fn list_models(&self) -> Result, LlmError> { // let models = models::ModelsDevService::new(&self.redis, &self.client) // .list_models({ diff --git a/server-new/src/llm/types.rs b/server-new/src/llm/types.rs index 9075b06..5f6e8fc 100644 --- a/server-new/src/llm/types.rs +++ b/server-new/src/llm/types.rs @@ -1,6 +1,12 @@ use serde::{Deserialize, Serialize}; -/// Generic chat request for all LLM providers +/// Generic LLM prompt +pub struct LlmPrompt<'r> { + pub text: &'r str, + pub options: &'r LlmChatOptions, +} + +/// Generic LLM chat request pub struct LlmChatRequest<'r> { pub messages: &'r [LlmMessage], // tools: Option>, diff --git a/server-new/src/plugins/auth.rs b/server-new/src/plugins/auth.rs index 773fe9c..ffec250 100644 --- a/server-new/src/plugins/auth.rs +++ b/server-new/src/plugins/auth.rs @@ -9,7 +9,8 @@ use crate::{ config::AppConfig, db::DbPool, services::auth::{ - oauth::OAuthService, session::AuthSessionService, session_store::SessionDbStore, + encryption::Encryptor, oauth::OAuthService, session::AuthSessionService, + session_store::SessionDbStore, }, state::AppState, }; @@ -24,10 +25,21 @@ pub fn plugin() -> AdHocPlugin { let config = state.get::().context("no config")?; let db_pool = state.get::().context("no db pool")?.to_owned(); + // Verify encryption key and build encryptor + let encryption_key = hex::decode(&config.auth.encryption_key) + .context("encryption_key must be hex value")?; + if encryption_key.len() != 32 { + bail!("encryption_key must be 32 bytes"); + } + let encryptor = Encryptor::new(&encryption_key)?; + // Build configured OAuth providers - state.insert(OAuthService::build_provider_map(&config.auth)); + let oauth_providers = OAuthService::build_provider_map(&config.auth); + + state.insert(encryptor); + state.insert(oauth_providers); - // Session cleanup task + // Start session cleanup task tokio::spawn(async move { let mut interval = tokio::time::interval(CLEANUP_INTERVAL); interval.tick().await; @@ -44,12 +56,6 @@ pub fn plugin() -> AdHocPlugin { Ok(state) }) .on_setup(|router, state: &AppState| { - let cookie_key = hex::decode(&state.config.auth.cookie_key) - .context("cookie_key must be hex value")?; - if cookie_key.len() < 32 { - bail!("cookie_key must be at least 32 bytes"); - } - // Session persistence let redis_store = RedisStore::with_prefix(state.redis.clone(), REDIS_PREFIX.to_owned()); let db_store = SessionDbStore::new(state.db_pool.clone()); @@ -59,7 +65,9 @@ pub fn plugin() -> AdHocPlugin { let session_layer = SessionManagerLayer::new(session_store) .with_name(state.config.auth.cookie_name.clone()) .with_expiry(Expiry::OnInactivity(cookie::time::Duration::minutes(15))) // default short session for login/OAuth - .with_private(cookie::Key::derive_from(&cookie_key)) + .with_private(cookie::Key::derive_from(&hex::decode( + &state.config.auth.encryption_key, + )?)) .with_path("/") .with_secure(true) .with_http_only(true) diff --git a/server-new/src/services/auth/encryption.rs b/server-new/src/services/auth/encryption.rs new file mode 100644 index 0000000..0163445 --- /dev/null +++ b/server-new/src/services/auth/encryption.rs @@ -0,0 +1,76 @@ +use aes_gcm::{ + Aes256Gcm, Nonce, + aead::{Aead, Generate, KeyInit}, +}; + +/// Service for encrypting and decrypting secrets +pub struct Encryptor { + cipher: Aes256Gcm, +} + +type EncryptorResult = Result; + +/// Errors that can occur during encryption / decryption +#[derive(Debug, thiserror::Error)] +pub enum EncryptorError { + #[error("encryption error")] + Encryption, + #[error("decryption error")] + Decryption, + #[error("invalid key")] + InvalidKey, + #[error("invalid nonce")] + InvalidNonce, +} + +impl Encryptor { + pub fn new(key_bytes: &[u8]) -> EncryptorResult { + let cipher = + Aes256Gcm::new_from_slice(key_bytes).map_err(|_| EncryptorError::InvalidKey)?; + Ok(Self { cipher }) + } + + /// Encrypts a string using AES-256-GCM and returns the ciphertext and nonce. + pub fn encrypt_string(&self, plaintext: &str) -> EncryptorResult<(Vec, Vec)> { + let nonce = Nonce::generate(); + let ciphertext = self + .cipher + .encrypt(&nonce, plaintext.as_bytes()) + .map_err(|_| EncryptorError::Encryption)?; + + Ok((ciphertext, nonce.to_vec())) + } + + /// Encrypts a byte slice using AES-256-GCM and returns the ciphertext and nonce. + pub fn encrypt_bytes(&self, bytes: &[u8]) -> EncryptorResult<(Vec, Vec)> { + let nonce = Nonce::generate(); + let ciphertext = self + .cipher + .encrypt(&nonce, bytes) + .map_err(|_| EncryptorError::Encryption)?; + + Ok((ciphertext, nonce.to_vec())) + } + + /// Decrypts a string using AES-256-GCM. + pub fn decrypt_string(&self, ciphertext: &[u8], nonce: &[u8]) -> EncryptorResult { + let nonce = Nonce::try_from(nonce).map_err(|_| EncryptorError::InvalidNonce)?; + let plaintext = self + .cipher + .decrypt(&nonce, ciphertext) + .map_err(|_| EncryptorError::Decryption)?; + + Ok(String::from_utf8(plaintext).map_err(|_| EncryptorError::Decryption)?) + } + + /// Decrypts a byte slice using AES-256-GCM. + pub fn decrypt_bytes(&self, ciphertext: &[u8], nonce: &[u8]) -> EncryptorResult> { + let nonce = Nonce::try_from(nonce).map_err(|_| EncryptorError::InvalidNonce)?; + let bytes = self + .cipher + .decrypt(&nonce, ciphertext) + .map_err(|_| EncryptorError::Decryption)?; + + Ok(bytes) + } +} diff --git a/server-new/src/services/auth/error.rs b/server-new/src/services/auth/error.rs index d03b019..aa9ca27 100644 --- a/server-new/src/services/auth/error.rs +++ b/server-new/src/services/auth/error.rs @@ -1,4 +1,4 @@ -use crate::{db::DbPoolError, error::AppError}; +use crate::error::AppError; pub type AuthResult = Result; @@ -16,7 +16,7 @@ pub enum AuthError { #[error("database error: {0}")] Database(#[from] diesel::result::Error), #[error("database pool error: {0}")] - DatabasePool(#[from] DbPoolError), + DatabasePool(#[from] crate::db::DbPoolError), #[error("session error: {0}")] Session(#[from] tower_sessions::session::Error), } diff --git a/server-new/src/services/auth/mod.rs b/server-new/src/services/auth/mod.rs index 6ae5d2b..fbfe8f5 100644 --- a/server-new/src/services/auth/mod.rs +++ b/server-new/src/services/auth/mod.rs @@ -1,20 +1,18 @@ use crate::{ config::AppConfig, - db::{DbPool, DbService, models::ChatRsUser}, + db::{DbService, models::ChatRsUser}, }; use uuid::Uuid; +pub mod encryption; mod error; pub mod oauth; pub mod session; pub mod session_store; -mod types; -pub use error::{AuthError, AuthResult}; -pub use types::*; +use error::{AuthError, AuthResult}; pub struct AuthService<'a> { - db: &'a DbPool, config: &'a AppConfig, http_client: &'a reqwest::Client, oauth_providers: &'a oauth::OAuthProviderMap, @@ -22,13 +20,11 @@ pub struct AuthService<'a> { impl<'a> AuthService<'a> { pub fn new( - db: &'a DbPool, config: &'a AppConfig, http_client: &'a reqwest::Client, oauth_providers: &'a oauth::OAuthProviderMap, ) -> Self { Self { - db, config, http_client, oauth_providers, @@ -37,8 +33,7 @@ impl<'a> AuthService<'a> { /// Get the user from the database with the given ID, or return /// an internal error if not found - pub async fn get_user(&self, id: &Uuid) -> AuthResult { - let mut db = DbService::from_pool(&self.db).await?; + pub async fn get_user(&self, db: &mut DbService, id: &Uuid) -> AuthResult { match db.users().find_by_id(id).await? { None => Err(AuthError::UserNotFound), Some(user) => Ok(user), @@ -52,6 +47,6 @@ impl<'a> AuthService<'a> { /// Access OAuth functions pub fn oauth(self) -> oauth::OAuthService<'a> { - oauth::OAuthService::new(self.config, self.db, self.http_client, self.oauth_providers) + oauth::OAuthService::new(self.config, self.http_client, self.oauth_providers) } } diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index fd2fcbe..6a9f676 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -11,10 +11,11 @@ use tower_sessions::Session; use crate::{ config::AppConfig, db::{ - DbPool, DbService, + DbService, models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, }, - services::auth::{AuthError, AuthResult, UserSession}, + extractors::session::UserSession, + services::auth::{AuthError, AuthResult}, }; mod discord; @@ -67,7 +68,6 @@ pub trait OAuthProvider: Send + Sync { /// OAuth functions pub struct OAuthService<'a> { config: &'a AppConfig, - db: &'a DbPool, http_client: &'a reqwest::Client, provider_map: &'a OAuthProviderMap, } @@ -78,13 +78,11 @@ impl<'a> OAuthService<'a> { pub(super) fn new( config: &'a AppConfig, - db: &'a DbPool, http_client: &'a reqwest::Client, provider_map: &'a OAuthProviderMap, ) -> Self { Self { config, - db, http_client, provider_map, } @@ -147,6 +145,7 @@ impl<'a> OAuthService<'a> { pub async fn get_user( &self, + db: &mut DbService, provider: OAuthProviderEnum, token: &StandardTokenResponse, active_session: Option, @@ -159,8 +158,7 @@ impl<'a> OAuthService<'a> { .await?; // Check for existing user, or create new user - let mut db = DbService::from_pool(self.db).await?; - let user = match oauth_provider.find_linked_user(&mut db, &user_info).await? { + let user = match oauth_provider.find_linked_user(db, &user_info).await? { Some(existing_user) => { if active_session.is_some_and(|sess| sess.user_id != existing_user.id) { return Err(AuthError::Unauthorized("cannot switch users via OAuth")); diff --git a/server-new/src/services/auth/session.rs b/server-new/src/services/auth/session.rs index 3ff3102..35ce6cb 100644 --- a/server-new/src/services/auth/session.rs +++ b/server-new/src/services/auth/session.rs @@ -9,10 +9,8 @@ use uuid::Uuid; use crate::{ db::{DbPool, DbService}, - services::auth::{ - AuthResult, - types::{SessionMeta, UserSession}, - }, + extractors::session::{SessionMeta, UserSession}, + services::auth::AuthResult, }; /// The field used to store the user ID in the session. @@ -62,7 +60,7 @@ impl AuthSessionService { #[tracing::instrument(skip(db_pool), level = "debug")] pub async fn session_cleanup(db_pool: &DbPool) -> AuthResult { let mut db = DbService::from_pool(&db_pool).await?; - Ok(db.sessions().delete_expired().await?) + Ok(db.auth_sessions().delete_expired().await?) } } diff --git a/server-new/src/services/auth/session_store.rs b/server-new/src/services/auth/session_store.rs index 8608be4..8237e1d 100644 --- a/server-new/src/services/auth/session_store.rs +++ b/server-new/src/services/auth/session_store.rs @@ -50,7 +50,7 @@ impl SessionStore for SessionDbStore { let expires_at = Self::convert_expiry(record.expiry_date)?; let mut db = self.get_db().await?; - db.sessions() + db.auth_sessions() .create(&session_id, user_id.as_ref(), &record.data, expires_at) .await .map_err(|err| Error::Backend(err.to_string()))?; @@ -66,7 +66,7 @@ impl SessionStore for SessionDbStore { let expires_at = Self::convert_expiry(record.expiry_date)?; let mut db = self.get_db().await?; - db.sessions() + db.auth_sessions() .update(&session_id, &record.data, expires_at) .await .map_err(|err| Error::Backend(err.to_string()))?; @@ -83,7 +83,7 @@ impl SessionStore for SessionDbStore { let session_id = Self::get_session_uuid(&session_id); let mut db = self.get_db().await?; - match db.sessions().find_active_by_id(&session_id).await { + match db.auth_sessions().find_active_by_id(&session_id).await { Ok(Some(session)) => Ok(Some(Record { id: Id(i128::from_be_bytes(session.id.into_bytes())), data: session.data.0, @@ -102,7 +102,7 @@ impl SessionStore for SessionDbStore { let session_id = Self::get_session_uuid(session_id); let mut db = self.get_db().await?; - db.sessions() + db.auth_sessions() .delete_by_id(&session_id) .await .map_err(|err| Error::Backend(err.to_string()))?; diff --git a/server-new/src/services/auth/types.rs b/server-new/src/services/auth/types.rs deleted file mode 100644 index c9206f6..0000000 --- a/server-new/src/services/auth/types.rs +++ /dev/null @@ -1,26 +0,0 @@ -use std::net::IpAddr; - -use serde::{Deserialize, Serialize}; -use uuid::Uuid; - -use crate::db::UtcDateTime; - -/// Active user session data. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct UserSession { - pub user_id: Uuid, -} - -impl UserSession { - pub fn new(user_id: Uuid) -> Self { - Self { user_id } - } -} - -/// Session metadata captured on login. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SessionMeta { - pub start_time: UtcDateTime, - pub ip: Option, - pub user_agent: Option, -} diff --git a/server-new/src/services/chat/error.rs b/server-new/src/services/chat/error.rs index dd45481..48253dc 100644 --- a/server-new/src/services/chat/error.rs +++ b/server-new/src/services/chat/error.rs @@ -1,10 +1,16 @@ +use crate::error::AppError; + /// Chat service errors #[derive(Debug, thiserror::Error)] pub enum ChatError { + #[error("invalid message history")] + Messages, + #[error("session not found")] + SessionNotFound, + #[error("already streaming a response")] + AlreadyStreaming, #[error(transparent)] Request(#[from] crate::llm::error::LlmRequestError), - #[error("Invalid message history")] - Messages, #[error(transparent)] Streaming(#[from] crate::services::stream::error::StreamingError), #[error("database error: {0}")] @@ -12,3 +18,15 @@ pub enum ChatError { #[error("database pool error: {0}")] DatabasePool(#[from] crate::db::DbPoolError), } + +impl From for AppError { + fn from(value: ChatError) -> Self { + match value { + ChatError::Messages => Self::bad_request("invalid messages"), + ChatError::SessionNotFound => Self::not_found("chat session not found"), + ChatError::AlreadyStreaming => Self::bad_request("already streaming this chat session"), + ChatError::Request(err) => Self::bad_request(err.to_string()), + err => Self::internal(err.into()), + } + } +} diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index cff77d7..9530bf0 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -1,3 +1,5 @@ +use std::sync::Arc; + use tinistream_client::types::StreamAccessResponse; use uuid::Uuid; @@ -21,6 +23,9 @@ use crate::{ mod error; mod messages; +mod titles; + +const DEFAULT_SESSION_TITLE: &str = "New Chat"; pub struct ChatService<'r> { db_pool: &'r DbPool, @@ -37,44 +42,64 @@ impl<'r> ChatService<'r> { pub async fn stream_user_chat( &self, + db: &mut DbService, user_id: Uuid, session_id: Uuid, - provider: &dyn LlmProvider, provider_id: i32, - user_message: LlmUserMessage, + provider: Arc, + user_message: Option, chat_options: LlmChatOptions, ) -> Result { - // Check for existing chat session, then save the new user message to it - let mut db = DbService::from_pool(self.db_pool).await?; - let (_existing_session, mut session_messages) = db - .chats() - .get_session_with_messages(&user_id, &session_id) - .await?; - let new_message = db + let stream_key = StreamingService::chat_stream_key(&user_id, &session_id); + let stream_service = StreamingService::new(self.tinistream); + + // Get session and message history + let (chat_session, mut session_messages) = db .chats() - .save_message(NewChatRsMessage { - session_id: &session_id, - role: ChatRsMessageRole::User, - content: &user_message.text, - meta: ChatRsMessageMeta::new_user(UserMeta::default()), - }) + .find_session_with_messages(&user_id, &session_id) .await?; - session_messages.push(new_message); + let chat_session = chat_session.ok_or(ChatError::SessionNotFound)?; + + // Check that we're not already streaming a response for this chat session + if stream_service.exists_stream(&stream_key).await? { + return Err(ChatError::AlreadyStreaming); + } + + // Save user message, and generate session title if needed + if let Some(user_message) = user_message { + if session_messages.is_empty() && chat_session.title == DEFAULT_SESSION_TITLE { + titles::generate_title( + user_id, + session_id, + &user_message.text, + &provider, + &chat_options.model, + self.db_pool, + ); + } + let new_message = db + .chats() + .save_message(NewChatRsMessage { + content: &user_message.text, + session_id: &session_id, + role: ChatRsMessageRole::User, + meta: ChatRsMessageMeta::new_user(UserMeta::default()), + }) + .await?; + session_messages.push(new_message); + } // Send the request to the LLM provider and get the streaming response - let llm_messages = messages::build_llm_messages(session_messages)?; let stream = provider .stream_chat(LlmChatRequest { - messages: &llm_messages, + messages: &messages::build_llm_messages(session_messages)?, options: &chat_options, }) .await?; // Create a new client stream in `tinistream` to stream the response to the user - let stream_key = StreamingService::chat_stream_key(&user_id, &session_id); - let (stream_access, ws_writer, ws_reader) = StreamingService::new(self.tinistream) - .create_stream(&stream_key) - .await?; + let (stream_access, ws_writer, ws_reader) = + stream_service.create_stream(&stream_key).await?; // Spawn task to process and save the streaming response let db_pool = self.db_pool.to_owned(); @@ -88,6 +113,7 @@ impl<'r> ChatService<'r> { { tracing::error!("Failed to save assistant response: {err}"); } + if !stream_cancelled { let _ = StreamingService::new(&tinistream_client) .end_stream(&stream_key) diff --git a/server-new/src/services/chat/titles.rs b/server-new/src/services/chat/titles.rs new file mode 100644 index 0000000..12fafa8 --- /dev/null +++ b/server-new/src/services/chat/titles.rs @@ -0,0 +1,74 @@ +use std::sync::Arc; + +use uuid::Uuid; + +use crate::{ + db::{DbPool, DbService, models::UpdateChatRsSession}, + llm::{ + interface::LlmProvider, + types::{LlmChatOptions, LlmPrompt}, + }, + services::chat::error::ChatError, +}; + +/// Spawn a task to generate a title for the chat session +pub fn generate_title( + user_id: Uuid, + session_id: Uuid, + first_message: &str, + provider: &Arc, + model: &str, + pool: &DbPool, +) { + let user_message = first_message.to_owned(); + let provider = Arc::clone(provider); + let model = model.to_owned(); + let pool = pool.to_owned(); + + tokio::spawn(async move { + if let Err(err) = generate(user_id, session_id, user_message, provider, model, pool).await { + tracing::warn!("Failed to generate title: {}", err); + } + }); +} + +const TITLE_PROMPT: &str = "This is the first message sent by a human in a chat session with an AI chatbot. \ + Please generate a short title for the chat session (3-7 words) in plain text, with no quotes or prefixes"; +const TITLE_PROMPT_TEMPERATURE: f32 = 0.7; +const TITLE_PROMPT_MAX_TOKENS: u32 = 20; + +async fn generate( + user_id: Uuid, + session_id: Uuid, + user_message: String, + provider: Arc, + model: String, + db_pool: DbPool, +) -> Result<(), ChatError> { + let message = format!("{TITLE_PROMPT}: \"{user_message}\""); + let title = provider + .prompt(LlmPrompt { + text: &message, + options: &LlmChatOptions { + model, + temperature: Some(TITLE_PROMPT_TEMPERATURE), + max_tokens: Some(TITLE_PROMPT_MAX_TOKENS), + ..Default::default() + }, + }) + .await?; + + let mut db = DbService::from_pool(&db_pool).await?; + db.chats() + .update_session( + &user_id, + &session_id, + UpdateChatRsSession { + title: Some(title.trim()), + ..Default::default() + }, + ) + .await?; + + Ok(()) +} diff --git a/server-new/src/services/mod.rs b/server-new/src/services/mod.rs index d2d9b1b..1e9f048 100644 --- a/server-new/src/services/mod.rs +++ b/server-new/src/services/mod.rs @@ -1,3 +1,4 @@ pub mod auth; pub mod chat; +pub mod provider; pub mod stream; diff --git a/server-new/src/services/provider/error.rs b/server-new/src/services/provider/error.rs new file mode 100644 index 0000000..80202f5 --- /dev/null +++ b/server-new/src/services/provider/error.rs @@ -0,0 +1,27 @@ +use crate::{ + db::models::ParseProviderTypeError, error::AppError, services::auth::encryption::EncryptorError, +}; + +#[derive(Debug, thiserror::Error)] +pub enum ProviderError { + #[error("provider not found")] + NotFound, + #[error("missing API key")] + MissingApiKey, + #[error(transparent)] + InvalidProviderType(#[from] ParseProviderTypeError), + #[error("error reading/writing API keys: {0}")] + Encryption(#[from] EncryptorError), + #[error("database error: {0}")] + Database(#[from] diesel::result::Error), +} + +impl From for AppError { + fn from(value: ProviderError) -> Self { + match value { + ProviderError::NotFound => Self::not_found("provider not found"), + ProviderError::MissingApiKey => Self::bad_request("missing API key for this provider"), + error => Self::internal(error.into()), + } + } +} diff --git a/server-new/src/services/provider/mod.rs b/server-new/src/services/provider/mod.rs new file mode 100644 index 0000000..751d302 --- /dev/null +++ b/server-new/src/services/provider/mod.rs @@ -0,0 +1,80 @@ +use std::{str::FromStr, sync::Arc}; + +use uuid::Uuid; + +use crate::{ + db::{ + DbService, + models::{ChatRsProvider, ChatRsProviderType, ChatRsSecret}, + }, + llm::{interface::LlmProvider, providers::OpenAIProvider}, + services::{auth::encryption::Encryptor, provider::error::ProviderError}, +}; + +mod error; + +pub struct ProviderService<'r> { + encryptor: &'r Encryptor, + http_client: &'r reqwest::Client, +} + +impl<'r> ProviderService<'r> { + pub fn new(encryptor: &'r Encryptor, http_client: &'r reqwest::Client) -> Self { + Self { + encryptor, + http_client, + } + } + + pub async fn get_provider( + &self, + db: &mut DbService, + user_id: &Uuid, + provider_id: i32, + ) -> Result<(ChatRsProvider, ChatRsProviderType, Option), ProviderError> { + let (provider, api_key_secret) = db + .providers() + .find_by_id(user_id, provider_id) + .await? + .ok_or(ProviderError::NotFound)?; + let provider_type = ChatRsProviderType::from_str(provider.provider_type.as_str())?; + + Ok((provider, provider_type, api_key_secret)) + } + + pub async fn build_llm_provider( + &self, + db: &mut DbService, + user_id: &Uuid, + provider_id: i32, + ) -> Result, ProviderError> { + let (_provider, provider_type, api_key_secret) = + self.get_provider(db, user_id, provider_id).await?; + let api_key = api_key_secret + .map(|secret| { + self.encryptor + .decrypt_string(&secret.ciphertext, &secret.nonce) + }) + .transpose()?; + + let llm_provider = match provider_type { + ChatRsProviderType::Openai => Arc::new(OpenAIProvider::openai( + self.http_client, + api_key.ok_or(ProviderError::MissingApiKey)?, + )), + _ => todo!(), + // ChatRsProviderType::Anthropic => Box::new(AnthropicProvider::new( + // http_client, + // redis, + // api_key.ok_or(ProviderError::MissingApiKey)?, + // )), + // ChatRsProviderType::Ollama => Box::new(OllamaProvider::new( + // http_client, + // base_url.unwrap_or("http://localhost:11434"), + // )), + // ChatRsProviderType::Lorem => Box::new(LoremProvider::new()), + }; + + Ok(llm_provider) + } +} diff --git a/server-new/src/services/stream/error.rs b/server-new/src/services/stream/error.rs index fb41cc5..2cc3b93 100644 --- a/server-new/src/services/stream/error.rs +++ b/server-new/src/services/stream/error.rs @@ -1,4 +1,4 @@ -/// Errors that can occur during streaming +/// Streaming infrastructure errors #[derive(Debug, thiserror::Error)] pub enum StreamingError { #[error("Client streaming error: {0}")] diff --git a/server-new/src/services/stream/mod.rs b/server-new/src/services/stream/mod.rs index b4d4324..4a04a98 100644 --- a/server-new/src/services/stream/mod.rs +++ b/server-new/src/services/stream/mod.rs @@ -52,6 +52,11 @@ impl<'r> StreamingService<'r> { format!("user:{}:chat:", user_id) } + /// Check for existing client stream in `tinistream` + pub async fn exists_stream(&self, stream_key: &str) -> Result { + Ok(self.tinistream.stream_exists(&stream_key).await?) + } + /// Start the client stream in `tinistream`, and return a WebSocket writer and reader for it pub async fn create_stream( &self, diff --git a/server-new/src/services/stream/writer.rs b/server-new/src/services/stream/writer.rs index bde6c77..341098c 100644 --- a/server-new/src/services/stream/writer.rs +++ b/server-new/src/services/stream/writer.rs @@ -50,8 +50,8 @@ struct ChunkState { #[serde(tag = "event", content = "data", rename_all = "snake_case")] pub(super) enum RedisStreamChunk { Text(String), - ToolCall(String), - PendingToolCall(String), + // ToolCall(String), + // PendingToolCall(String), Error(String), } diff --git a/server-new/src/state.rs b/server-new/src/state.rs index 2ea3b92..e7e24dc 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -8,8 +8,9 @@ use crate::{ config::AppConfig, db::DbPool, services::{ - auth::{AuthService, oauth::OAuthProviderMap}, + auth::{AuthService, encryption::Encryptor, oauth::OAuthProviderMap}, chat::ChatService, + provider::ProviderService, stream::tinistream::TinistreamClient, }, }; @@ -21,26 +22,24 @@ pub struct AppState(Arc); #[derive(AppState)] pub struct AppStateInner { pub config: AppConfig, - pub http_client: reqwest::Client, pub db_pool: DbPool, + pub encryptor: Encryptor, + pub http_client: reqwest::Client, + pub oauth_providers: OAuthProviderMap, pub redis: fred::prelude::Pool, pub tinistream: TinistreamClient, - pub oauth_providers: OAuthProviderMap, } impl AppState { pub fn auth_service(&self) -> AuthService<'_> { - AuthService::new( - &self.db_pool, - &self.config, - &self.http_client, - &self.oauth_providers, - ) + AuthService::new(&self.config, &self.http_client, &self.oauth_providers) } - pub fn chat_service(&self) -> ChatService<'_> { ChatService::new(&self.db_pool, &self.tinistream) } + pub fn provider_service(&self) -> ProviderService<'_> { + ProviderService::new(&self.encryptor, &self.http_client) + } } impl Deref for AppState { From 65f91f9d691bc1a614143bf7e69d19944a62382d Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 1 Jul 2026 11:25:17 -0400 Subject: [PATCH 048/111] tweaks --- server-new/src/llm/providers/utils.rs | 7 +++++-- server-new/src/services/chat/mod.rs | 7 +++---- server-new/src/services/chat/titles.rs | 2 +- server-new/src/services/stream/error.rs | 4 ++-- server-new/src/services/stream/tinistream.rs | 2 +- server-new/src/services/stream/writer.rs | 2 +- 6 files changed, 13 insertions(+), 11 deletions(-) diff --git a/server-new/src/llm/providers/utils.rs b/server-new/src/llm/providers/utils.rs index 327f8f7..9d2afd1 100644 --- a/server-new/src/llm/providers/utils.rs +++ b/server-new/src/llm/providers/utils.rs @@ -10,6 +10,9 @@ use tokio_util::{ use crate::llm::error::LlmStreamChunkError; +/// Max allowed length of stream lines (5 KB) +const MAX_LINE_LEN: usize = 5 * 1024; + /// Create a data URI pub fn create_data_uri(content_type: &str, b64_string: &str) -> String { format!("data:{content_type};base64,{b64_string}") @@ -20,7 +23,7 @@ pub fn get_sse_events( response: reqwest::Response, ) -> impl Stream> { let stream_reader = StreamReader::new(response.bytes_stream().map_err(std::io::Error::other)); - let line_reader = FramedRead::new(stream_reader, LinesCodec::new()); + let line_reader = FramedRead::new(stream_reader, LinesCodec::new_with_max_length(MAX_LINE_LEN)); line_reader.filter_map(|line_result| { match line_result { @@ -46,7 +49,7 @@ pub fn get_json_events( response: reqwest::Response, ) -> impl Stream> { let stream_reader = StreamReader::new(response.bytes_stream().map_err(std::io::Error::other)); - let line_reader = FramedRead::new(stream_reader, LinesCodec::new()); + let line_reader = FramedRead::new(stream_reader, LinesCodec::new_with_max_length(MAX_LINE_LEN)); line_reader.map(|line_result| match line_result { Ok(line) => serde_json::from_str::(&line).map_err(LlmStreamChunkError::Parsing), Err(e) => Err(LlmStreamChunkError::Decoding(e)), diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index 9530bf0..b3ce9e5 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -50,9 +50,6 @@ impl<'r> ChatService<'r> { user_message: Option, chat_options: LlmChatOptions, ) -> Result { - let stream_key = StreamingService::chat_stream_key(&user_id, &session_id); - let stream_service = StreamingService::new(self.tinistream); - // Get session and message history let (chat_session, mut session_messages) = db .chats() @@ -61,6 +58,8 @@ impl<'r> ChatService<'r> { let chat_session = chat_session.ok_or(ChatError::SessionNotFound)?; // Check that we're not already streaming a response for this chat session + let stream_key = StreamingService::chat_stream_key(&user_id, &session_id); + let stream_service = StreamingService::new(self.tinistream); if stream_service.exists_stream(&stream_key).await? { return Err(ChatError::AlreadyStreaming); } @@ -89,7 +88,7 @@ impl<'r> ChatService<'r> { session_messages.push(new_message); } - // Send the request to the LLM provider and get the streaming response + // Send the request to the LLM provider and get the stream response let stream = provider .stream_chat(LlmChatRequest { messages: &messages::build_llm_messages(session_messages)?, diff --git a/server-new/src/services/chat/titles.rs b/server-new/src/services/chat/titles.rs index 12fafa8..63f4c3b 100644 --- a/server-new/src/services/chat/titles.rs +++ b/server-new/src/services/chat/titles.rs @@ -27,7 +27,7 @@ pub fn generate_title( tokio::spawn(async move { if let Err(err) = generate(user_id, session_id, user_message, provider, model, pool).await { - tracing::warn!("Failed to generate title: {}", err); + tracing::warn!("Failed to generate title: {err}"); } }); } diff --git a/server-new/src/services/stream/error.rs b/server-new/src/services/stream/error.rs index 2cc3b93..bb5c252 100644 --- a/server-new/src/services/stream/error.rs +++ b/server-new/src/services/stream/error.rs @@ -1,8 +1,8 @@ /// Streaming infrastructure errors #[derive(Debug, thiserror::Error)] pub enum StreamingError { - #[error("Client streaming error: {0}")] + #[error("tinistream error: {0}")] Tinistream(#[from] super::tinistream::TiniError), - #[error("Websocket error: {0}")] + #[error("websocket error: {0}")] Websocket(#[from] reqwest_websocket::Error), } diff --git a/server-new/src/services/stream/tinistream.rs b/server-new/src/services/stream/tinistream.rs index 5000ca6..72706cd 100644 --- a/server-new/src/services/stream/tinistream.rs +++ b/server-new/src/services/stream/tinistream.rs @@ -13,7 +13,7 @@ pub struct TinistreamClient { pub type TiniResult = Result; #[derive(Debug, thiserror::Error)] -#[error("{message}")] +#[error("{status} {message}")] pub struct TiniError { pub status: u16, pub code: String, diff --git a/server-new/src/services/stream/writer.rs b/server-new/src/services/stream/writer.rs index 341098c..704f39a 100644 --- a/server-new/src/services/stream/writer.rs +++ b/server-new/src/services/stream/writer.rs @@ -147,7 +147,7 @@ impl LlmStreamWriter { fn process_text(&mut self, text: &str) { self.current_chunk .text - .get_or_insert_with(|| String::with_capacity(MAX_CHUNK_SIZE)) + .get_or_insert_with(|| String::with_capacity(MAX_CHUNK_SIZE * 2)) .push_str(text); self.complete_text .get_or_insert_with(|| String::with_capacity(1024)) From d85e44f0ae3a21f1f924a9358ec157792fd7b03c Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 1 Jul 2026 13:36:10 -0400 Subject: [PATCH 049/111] add lorem llm provider, moar auth tweaks --- server-new/Cargo.lock | 2 +- server-new/Cargo.toml | 2 +- server-new/src/llm/providers/lorem.rs | 108 +++++++++++++++++++ server-new/src/llm/providers/mod.rs | 2 + server-new/src/llm/providers/openai/mod.rs | 60 ++++------- server-new/src/llm/providers/utils.rs | 30 +++++- server-new/src/plugins/auth.rs | 3 +- server-new/src/services/auth/mod.rs | 10 +- server-new/src/services/auth/oauth.rs | 116 +++++++++++---------- server-new/src/services/provider/mod.rs | 9 +- server-new/src/state.rs | 2 +- 11 files changed, 232 insertions(+), 112 deletions(-) create mode 100644 server-new/src/llm/providers/lorem.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index dcecdfe..77a7d79 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -2597,7 +2597,7 @@ checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" [[package]] name = "simple-oauth" version = "0.1.0" -source = "git+https://github.com/fa-sharp/simple-oauth-rs?rev=45ab590#45ab590710c4830d509c8b5b1781fd19d381525b" +source = "git+https://github.com/fa-sharp/simple-oauth-rs?rev=9eebc3e#9eebc3ec19dddf83643e2796b14336bb92ccedd0" dependencies = [ "bon", "oauth2", diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index ea6819d..793f6a0 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -57,7 +57,7 @@ serde_with = { } simple-oauth = { git = "https://github.com/fa-sharp/simple-oauth-rs", - rev = "45ab590", + rev = "9eebc3e", features = ["default-tls"] } thiserror = "2.0.18" diff --git a/server-new/src/llm/providers/lorem.rs b/server-new/src/llm/providers/lorem.rs new file mode 100644 index 0000000..4b2e35f --- /dev/null +++ b/server-new/src/llm/providers/lorem.rs @@ -0,0 +1,108 @@ +//! Lorem ipsum LLM provider (for testing) + +use std::{pin::Pin, time::Duration}; + +use futures::Stream; +use tokio::time::{Interval, interval}; + +use crate::llm::{ + error::LlmStreamChunkError, + interface::{ + LlmPromptResponse, LlmProvider, LlmStream, LlmStreamChunk, LlmStreamChunkResult, + LlmStreamingResponse, + }, + types::{LlmChatRequest, LlmPrompt}, +}; + +/// A test/dummy provider that streams 'lorem ipsum...' and emits test errors during the stream +#[derive(Debug, Clone)] +pub struct LoremProvider { + pub interval: u32, +} + +impl LoremProvider { + pub fn new() -> Self { + LoremProvider { interval: 400 } + } +} + +struct LoremStream { + words: Vec<&'static str>, + index: usize, + interval: Interval, +} +impl Stream for LoremStream { + type Item = LlmStreamChunkResult; + + fn poll_next( + mut self: Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + if self.index >= self.words.len() { + return std::task::Poll::Ready(None); + } + + match Pin::new(&mut self.interval).poll_tick(cx) { + std::task::Poll::Ready(_) => { + let word = self.words[self.index]; + self.index += 1; + if self.index == 0 || self.index % 10 != 0 { + std::task::Poll::Ready(Some(Ok(LlmStreamChunk::Text(word.to_owned())))) + } else { + std::task::Poll::Ready(Some(Err(LlmStreamChunkError::Provider( + "Test error".into(), + )))) + } + } + std::task::Poll::Pending => std::task::Poll::Pending, + } + } +} + +impl LlmProvider for LoremProvider { + fn prompt<'r>(&'r self, _prompt: LlmPrompt<'r>) -> LlmPromptResponse<'r> { + Box::pin(async { Ok("Lorem ipsum".to_owned()) }) + } + + fn stream_chat<'r>(&'r self, _request: LlmChatRequest<'r>) -> LlmStreamingResponse<'r> { + let lorem_words = vec![ + "Lorem ipsum ", + "dolor sit ", + "amet, consectetur ", + "adipiscing elit, ", + "sed do", + " eiusmod tempor", + " incididunt ut", + " labore et", + " dolore magna ", + "aliqua. Ut ", + "enim ad ", + "minim veniam,", + " quis nostrud", + " exercitation ullamco", + " laboris nisi ", + "ut aliquip ", + "ex ea ", + "commodo consequat. ", + "Duis aute ", + "irure dolor ", + "in reprehenderit ", + "in voluptate ", + "velit esse ", + "cillum dolore ", + "eu fugiat ", + "nulla pariatur.", + ]; + + Box::pin(async move { + let stream: LlmStream = Box::pin(LoremStream { + words: lorem_words, + index: 0, + interval: interval(Duration::from_millis(self.interval.into())), + }); + tokio::time::sleep(Duration::from_millis(1000)).await; // Simulate initial request latency + + Ok(stream) + }) + } +} diff --git a/server-new/src/llm/providers/mod.rs b/server-new/src/llm/providers/mod.rs index bb2961a..14a6989 100644 --- a/server-new/src/llm/providers/mod.rs +++ b/server-new/src/llm/providers/mod.rs @@ -1,4 +1,6 @@ +mod lorem; mod openai; mod utils; +pub use lorem::LoremProvider; pub use openai::{OpenAIProvider, OpenAIProviderConfig, OpenAIProviderFlavor}; diff --git a/server-new/src/llm/providers/openai/mod.rs b/server-new/src/llm/providers/openai/mod.rs index 7831818..2b86d0f 100644 --- a/server-new/src/llm/providers/openai/mod.rs +++ b/server-new/src/llm/providers/openai/mod.rs @@ -5,7 +5,7 @@ use futures::StreamExt; use crate::llm::{ error::LlmRequestError, interface::{LlmPromptResponse, LlmProvider, LlmStreamingResponse}, - providers::utils, + providers::utils::{self, llm_api_request}, types::{LlmChatRequest, LlmPrompt, LlmUsage}, }; @@ -159,34 +159,25 @@ impl LlmProvider for OpenAIProvider { Box::pin(async move { let provider_name = self.config.flavor.name(); - let response = self - .client - .post(format!("{}/chat/completions", self.config.base_url)) - .bearer_auth(&self.config.api_key) - .json(&request) - .send() - .await - .map_err(|err| { - LlmRequestError::Provider(format!("{provider_name} request failed: {err}")) - })?; - if !response.status().is_success() { - let status = response.status(); - let error_text = response.text().await.unwrap_or_default(); - return Err(LlmRequestError::Provider(format!( - "{provider_name} API error {status}: {error_text}", - ))); - } - - let mut openai_response: OpenAIResponse = response.json().await.map_err(|err| { + let response = llm_api_request( + &self.client, + provider_name, + &format!("{}/chat/completions", self.config.base_url), + &self.config.api_key, + &request, + ) + .await?; + let mut response: OpenAIResponse = response.json().await.map_err(|err| { LlmRequestError::Provider(format!("Failed to parse response: {err}")) })?; - let text = openai_response + + let text = response .choices .get_mut(0) .and_then(|choice| choice.message.as_mut()) .and_then(|message| message.content.take()) .ok_or(LlmRequestError::NoContent)?; - if let Some(usage) = openai_response.usage { + if let Some(usage) = response.usage { let usage: LlmUsage = usage.into(); tracing::info!("Prompt usage: {usage:?}"); } @@ -216,23 +207,14 @@ impl LlmProvider for OpenAIProvider { let provider_name = self.config.flavor.name(); Box::pin(async move { - let response = self - .client - .post(format!("{}/chat/completions", self.config.base_url)) - .bearer_auth(&self.config.api_key) - .json(&request) - .send() - .await - .map_err(|e| { - LlmRequestError::Provider(format!("{provider_name} request failed: {e}")) - })?; - if !response.status().is_success() { - let status = response.status(); - let error_text = response.text().await.unwrap_or_default(); - return Err(LlmRequestError::Provider(format!( - "{provider_name} API error {status}: {error_text}", - ))); - } + let response = llm_api_request( + &self.client, + provider_name, + &format!("{}/chat/completions", self.config.base_url), + &self.config.api_key, + &request, + ) + .await?; let stream = async_stream::stream! { let mut sse_event_stream = utils::get_sse_events(response); diff --git a/server-new/src/llm/providers/utils.rs b/server-new/src/llm/providers/utils.rs index 9d2afd1..23b32a7 100644 --- a/server-new/src/llm/providers/utils.rs +++ b/server-new/src/llm/providers/utils.rs @@ -1,14 +1,14 @@ //! Utilities for working with LLM requests and responses use futures::TryStreamExt; -use serde::de::DeserializeOwned; +use serde::{Serialize, de::DeserializeOwned}; use tokio_stream::{Stream, StreamExt}; use tokio_util::{ codec::{FramedRead, LinesCodec}, io::StreamReader, }; -use crate::llm::error::LlmStreamChunkError; +use crate::llm::error::{LlmRequestError, LlmStreamChunkError}; /// Max allowed length of stream lines (5 KB) const MAX_LINE_LEN: usize = 5 * 1024; @@ -55,3 +55,29 @@ pub fn get_json_events( Err(e) => Err(LlmStreamChunkError::Decoding(e)), }) } + +/// Convenience function to make an API request to an LLM provider +pub async fn llm_api_request( + client: &reqwest::Client, + provider_name: &str, + url: &str, + token: &str, + request: &Req, +) -> Result { + let response = client + .post(url) + .bearer_auth(token) + .json(&request) + .send() + .await + .map_err(|e| LlmRequestError::Provider(format!("{provider_name} request failed: {e}")))?; + if !response.status().is_success() { + let status = response.status(); + let error_text = response.text().await.unwrap_or_default(); + return Err(LlmRequestError::Provider(format!( + "{provider_name} API error status {status}: {error_text}", + ))); + } + + Ok(response) +} diff --git a/server-new/src/plugins/auth.rs b/server-new/src/plugins/auth.rs index ffec250..8a801c8 100644 --- a/server-new/src/plugins/auth.rs +++ b/server-new/src/plugins/auth.rs @@ -23,6 +23,7 @@ pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Auth") .on_init(async |mut state| { let config = state.get::().context("no config")?; + let http_client = state.get::().context("no HTTP client")?; let db_pool = state.get::().context("no db pool")?.to_owned(); // Verify encryption key and build encryptor @@ -34,7 +35,7 @@ pub fn plugin() -> AdHocPlugin { let encryptor = Encryptor::new(&encryption_key)?; // Build configured OAuth providers - let oauth_providers = OAuthService::build_provider_map(&config.auth); + let oauth_providers = OAuthService::build_provider_map(config, http_client); state.insert(encryptor); state.insert(oauth_providers); diff --git a/server-new/src/services/auth/mod.rs b/server-new/src/services/auth/mod.rs index fbfe8f5..b071a10 100644 --- a/server-new/src/services/auth/mod.rs +++ b/server-new/src/services/auth/mod.rs @@ -14,19 +14,13 @@ use error::{AuthError, AuthResult}; pub struct AuthService<'a> { config: &'a AppConfig, - http_client: &'a reqwest::Client, oauth_providers: &'a oauth::OAuthProviderMap, } impl<'a> AuthService<'a> { - pub fn new( - config: &'a AppConfig, - http_client: &'a reqwest::Client, - oauth_providers: &'a oauth::OAuthProviderMap, - ) -> Self { + pub fn new(config: &'a AppConfig, oauth_providers: &'a oauth::OAuthProviderMap) -> Self { Self { config, - http_client, oauth_providers, } } @@ -47,6 +41,6 @@ impl<'a> AuthService<'a> { /// Access OAuth functions pub fn oauth(self) -> oauth::OAuthService<'a> { - oauth::OAuthService::new(self.config, self.http_client, self.oauth_providers) + oauth::OAuthService::new(self.config, self.oauth_providers) } } diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index 6a9f676..c7896d4 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -3,7 +3,7 @@ use std::collections::HashMap; use futures::future::BoxFuture; use serde::{Deserialize, Serialize}; use simple_oauth::{ - SimpleOAuthClient, SimpleOAuthProvider, + SimpleOAuthClient, SimpleOAuthError, SimpleOAuthProvider, types::{OAuthCredentials, StandardTokenResponse, UserInfo}, }; use tower_sessions::Session; @@ -29,7 +29,9 @@ pub use google::GoogleOAuthConfig; pub use oidc::OidcConfig; /// Map of configured OAuth providers stored in state -pub type OAuthProviderMap = HashMap>; +pub type OAuthProviderMap = HashMap)>; +/// Type of the OAuth client stored in state +pub type OAuthClient = SimpleOAuthClient>; /// Supported OAuth providers #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] @@ -68,7 +70,6 @@ pub trait OAuthProvider: Send + Sync { /// OAuth functions pub struct OAuthService<'a> { config: &'a AppConfig, - http_client: &'a reqwest::Client, provider_map: &'a OAuthProviderMap, } @@ -76,27 +77,36 @@ impl<'a> OAuthService<'a> { const SESS_STATE_FIELD: &'static str = "oauth_state"; const SESS_PKCE_FIELD: &'static str = "oauth_verifier"; - pub(super) fn new( - config: &'a AppConfig, - http_client: &'a reqwest::Client, - provider_map: &'a OAuthProviderMap, - ) -> Self { + pub(super) fn new(config: &'a AppConfig, provider_map: &'a OAuthProviderMap) -> Self { Self { config, - http_client, provider_map, } } + fn oauth_provider( + &self, + provider: OAuthProviderEnum, + ) -> AuthResult<(&OAuthClient, &dyn OAuthProvider)> { + let (client, provider) = self + .provider_map + .get(&provider) + .ok_or_else(|| AuthError::BadRequest("unsupported OAuth provider"))?; + Ok((client, provider.as_ref())) + } + + fn get_redirect_url(&self, callback_path: &str) -> String { + format!("{}{}", &self.config.server.base_url, callback_path) + } + pub async fn authorize_url( &self, provider: OAuthProviderEnum, callback_path: &str, session: &Session, ) -> AuthResult { - let oauth_provider = self.oauth_provider(provider)?; - let auth = self - .oauth_client(oauth_provider)? + let (oauth_client, _) = self.oauth_provider(provider)?; + let auth = oauth_client .authorize_url() .redirect_url(self.get_redirect_url(callback_path)) .build()?; @@ -128,9 +138,8 @@ impl<'a> OAuthService<'a> { .ok_or(AuthError::Unauthorized("missing PKCE in session"))?; // Exchange code for token - let oauth_provider = self.oauth_provider(provider)?; - let response = self - .oauth_client(oauth_provider)? + let (oauth_client, _) = self.oauth_provider(provider)?; + let response = oauth_client .exchange_code() .redirect_url(self.get_redirect_url(callback_path)) .code(code) @@ -151,11 +160,8 @@ impl<'a> OAuthService<'a> { active_session: Option, ) -> AuthResult { // Get user info from provider - let oauth_provider = self.oauth_provider(provider)?; - let user_info = self - .oauth_client(oauth_provider)? - .get_user_info(&token.access_token) - .await?; + let (oauth_client, oauth_provider) = self.oauth_provider(provider)?; + let user_info = oauth_client.get_user_info(&token.access_token).await?; // Check for existing user, or create new user let user = match oauth_provider.find_linked_user(db, &user_info).await? { @@ -191,52 +197,50 @@ impl<'a> OAuthService<'a> { Ok(user) } - fn get_redirect_url(&self, callback_path: &str) -> String { - format!("{}{}", &self.config.server.base_url, callback_path) - } - - fn oauth_provider(&self, provider: OAuthProviderEnum) -> AuthResult<&dyn OAuthProvider> { - let provider = self - .provider_map - .get(&provider) - .ok_or_else(|| AuthError::BadRequest("unsupported OAuth provider"))?; - Ok(provider.as_ref()) - } - - fn oauth_client( - &self, - provider: &dyn OAuthProvider, - ) -> Result>, AuthError> { - Ok(simple_oauth::SimpleOAuthClient::builder() - .provider(provider.get_inner_provider()) - .credentials(provider.get_credentials()) - .http_client(self.http_client) - .build()?) - } - - pub fn build_provider_map(config: &crate::config::AuthConfig) -> OAuthProviderMap { + pub fn build_provider_map( + config: &crate::config::AppConfig, + http_client: &reqwest::Client, + ) -> Result { use { discord::DiscordProvider, github::GitHubProvider, google::GoogleProvider, oidc::OidcProvider, }; let mut map: OAuthProviderMap = HashMap::new(); - if let Some(ref c) = config.github { - map.insert(OAuthProviderEnum::Github, Box::new(GitHubProvider::new(c))); + if let Some(ref c) = config.auth.github { + let provider = Box::new(GitHubProvider::new(c)); + let client = Self::build_oauth_client(http_client, &*provider)?; + map.insert(OAuthProviderEnum::Github, (client, provider)); } - if let Some(ref c) = config.discord { - map.insert( - OAuthProviderEnum::Discord, - Box::new(DiscordProvider::new(c)), - ); + if let Some(ref c) = config.auth.discord { + let provider = Box::new(DiscordProvider::new(c)); + let client = Self::build_oauth_client(http_client, &*provider)?; + map.insert(OAuthProviderEnum::Discord, (client, provider)); } - if let Some(ref c) = config.google { - map.insert(OAuthProviderEnum::Google, Box::new(GoogleProvider::new(c))); + if let Some(ref c) = config.auth.google { + let provider = Box::new(GoogleProvider::new(c)); + let client = Self::build_oauth_client(http_client, &*provider)?; + map.insert(OAuthProviderEnum::Google, (client, provider)); } - if let Some(ref c) = config.oidc { - map.insert(OAuthProviderEnum::Oidc, Box::new(OidcProvider::new(c))); + if let Some(ref c) = config.auth.oidc { + let provider = Box::new(OidcProvider::new(c)); + let client = Self::build_oauth_client(http_client, &*provider)?; + map.insert(OAuthProviderEnum::Oidc, (client, provider)); } - map + Ok(map) + } + + fn build_oauth_client( + http_client: &reqwest::Client, + provider: &dyn OAuthProvider, + ) -> Result>, SimpleOAuthError> { + let oauth_client = SimpleOAuthClient::builder() + .credentials(provider.get_credentials()) + .redirect_url("http://example.com/should-be-overridden") + .provider(provider.get_inner_provider()) + .http_client(http_client) + .build()?; + Ok(oauth_client) } } diff --git a/server-new/src/services/provider/mod.rs b/server-new/src/services/provider/mod.rs index 751d302..a31a98d 100644 --- a/server-new/src/services/provider/mod.rs +++ b/server-new/src/services/provider/mod.rs @@ -7,7 +7,10 @@ use crate::{ DbService, models::{ChatRsProvider, ChatRsProviderType, ChatRsSecret}, }, - llm::{interface::LlmProvider, providers::OpenAIProvider}, + llm::{ + interface::LlmProvider, + providers::{LoremProvider, OpenAIProvider}, + }, services::{auth::encryption::Encryptor, provider::error::ProviderError}, }; @@ -57,7 +60,8 @@ impl<'r> ProviderService<'r> { }) .transpose()?; - let llm_provider = match provider_type { + let llm_provider: Arc = match provider_type { + ChatRsProviderType::Lorem => Arc::new(LoremProvider::new()), ChatRsProviderType::Openai => Arc::new(OpenAIProvider::openai( self.http_client, api_key.ok_or(ProviderError::MissingApiKey)?, @@ -72,7 +76,6 @@ impl<'r> ProviderService<'r> { // http_client, // base_url.unwrap_or("http://localhost:11434"), // )), - // ChatRsProviderType::Lorem => Box::new(LoremProvider::new()), }; Ok(llm_provider) diff --git a/server-new/src/state.rs b/server-new/src/state.rs index e7e24dc..c812e22 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -32,7 +32,7 @@ pub struct AppStateInner { impl AppState { pub fn auth_service(&self) -> AuthService<'_> { - AuthService::new(&self.config, &self.http_client, &self.oauth_providers) + AuthService::new(&self.config, &self.oauth_providers) } pub fn chat_service(&self) -> ChatService<'_> { ChatService::new(&self.db_pool, &self.tinistream) From a653b6003510dba5ef750b0429b70b818ede37fc Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 1 Jul 2026 13:45:00 -0400 Subject: [PATCH 050/111] Update oauth.rs --- server-new/src/services/auth/oauth.rs | 32 +++++++++++++-------------- 1 file changed, 16 insertions(+), 16 deletions(-) diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index c7896d4..4d38ef3 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -87,12 +87,12 @@ impl<'a> OAuthService<'a> { fn oauth_provider( &self, provider: OAuthProviderEnum, - ) -> AuthResult<(&OAuthClient, &dyn OAuthProvider)> { + ) -> AuthResult<(&OAuthClient, &Box)> { let (client, provider) = self .provider_map .get(&provider) .ok_or_else(|| AuthError::BadRequest("unsupported OAuth provider"))?; - Ok((client, provider.as_ref())) + Ok((client, provider)) } fn get_redirect_url(&self, callback_path: &str) -> String { @@ -208,24 +208,24 @@ impl<'a> OAuthService<'a> { let mut map: OAuthProviderMap = HashMap::new(); if let Some(ref c) = config.auth.github { - let provider = Box::new(GitHubProvider::new(c)); - let client = Self::build_oauth_client(http_client, &*provider)?; - map.insert(OAuthProviderEnum::Github, (client, provider)); + let provider = GitHubProvider::new(c); + let client = Self::build_oauth_client(http_client, &provider)?; + map.insert(OAuthProviderEnum::Github, (client, Box::new(provider))); } if let Some(ref c) = config.auth.discord { - let provider = Box::new(DiscordProvider::new(c)); - let client = Self::build_oauth_client(http_client, &*provider)?; - map.insert(OAuthProviderEnum::Discord, (client, provider)); + let provider = DiscordProvider::new(c); + let client = Self::build_oauth_client(http_client, &provider)?; + map.insert(OAuthProviderEnum::Discord, (client, Box::new(provider))); } if let Some(ref c) = config.auth.google { - let provider = Box::new(GoogleProvider::new(c)); - let client = Self::build_oauth_client(http_client, &*provider)?; - map.insert(OAuthProviderEnum::Google, (client, provider)); + let provider = GoogleProvider::new(c); + let client = Self::build_oauth_client(http_client, &provider)?; + map.insert(OAuthProviderEnum::Google, (client, Box::new(provider))); } if let Some(ref c) = config.auth.oidc { - let provider = Box::new(OidcProvider::new(c)); - let client = Self::build_oauth_client(http_client, &*provider)?; - map.insert(OAuthProviderEnum::Oidc, (client, provider)); + let provider = OidcProvider::new(c); + let client = Self::build_oauth_client(http_client, &provider)?; + map.insert(OAuthProviderEnum::Oidc, (client, Box::new(provider))); } Ok(map) @@ -233,12 +233,12 @@ impl<'a> OAuthService<'a> { fn build_oauth_client( http_client: &reqwest::Client, - provider: &dyn OAuthProvider, + provider: &impl OAuthProvider, ) -> Result>, SimpleOAuthError> { let oauth_client = SimpleOAuthClient::builder() + .provider(provider.get_inner_provider()) .credentials(provider.get_credentials()) .redirect_url("http://example.com/should-be-overridden") - .provider(provider.get_inner_provider()) .http_client(http_client) .build()?; Ok(oauth_client) From 24476b9f52bef6119d2e47649281ea9a39a18999 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 1 Jul 2026 13:59:01 -0400 Subject: [PATCH 051/111] improve session messages DB query --- server-new/src/db/repositories/chat.rs | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/server-new/src/db/repositories/chat.rs b/server-new/src/db/repositories/chat.rs index d7237af..f24e70b 100644 --- a/server-new/src/db/repositories/chat.rs +++ b/server-new/src/db/repositories/chat.rs @@ -127,6 +127,10 @@ impl<'a> ChatRepository<'a> { .select(ChatRsSession::as_select()) .first(&mut &**self.db), chat_messages::table + .inner_join( + chat_sessions::table.on(chat_sessions::id.eq(chat_messages::session_id)), + ) + .filter(chat_sessions::user_id.eq(user_id)) .filter(chat_messages::session_id.eq(session_id)) .select(ChatRsMessage::as_select()) .order_by(chat_messages::created_at.asc()) From 75588303561e6fb614d4ea5529b2a92334a89ac8 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 1 Jul 2026 14:03:14 -0400 Subject: [PATCH 052/111] fix OAuth setup --- server-new/src/plugins/auth.rs | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/server-new/src/plugins/auth.rs b/server-new/src/plugins/auth.rs index 8a801c8..999826b 100644 --- a/server-new/src/plugins/auth.rs +++ b/server-new/src/plugins/auth.rs @@ -35,7 +35,8 @@ pub fn plugin() -> AdHocPlugin { let encryptor = Encryptor::new(&encryption_key)?; // Build configured OAuth providers - let oauth_providers = OAuthService::build_provider_map(config, http_client); + let oauth_providers = OAuthService::build_provider_map(config, http_client) + .context("build OAuth providers")?; state.insert(encryptor); state.insert(oauth_providers); From 22ad88bf99c8800f328dd838cac42c79703fd26d Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 1 Jul 2026 15:15:07 -0400 Subject: [PATCH 053/111] add openapi with utoipa --- server-new/Cargo.lock | 60 +++++++++++++++++++++++++++ server-new/Cargo.toml | 3 ++ server-new/src/api/auth.rs | 20 +++++---- server-new/src/api/chat.rs | 21 +++++++--- server-new/src/api/health.rs | 7 +++- server-new/src/api/mod.rs | 28 ++++++++++--- server-new/src/db/models/user.rs | 7 ++-- server-new/src/services/auth/oauth.rs | 3 +- 8 files changed, 123 insertions(+), 26 deletions(-) diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index 77a7d79..e714170 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -1411,6 +1411,8 @@ checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" dependencies = [ "equivalent", "hashbrown", + "serde", + "serde_core", ] [[package]] @@ -1774,6 +1776,12 @@ dependencies = [ "windows-link", ] +[[package]] +name = "paste" +version = "1.0.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" + [[package]] name = "pear" version = "0.2.9" @@ -2267,6 +2275,9 @@ dependencies = [ "tracing", "tracing-appender", "tracing-subscriber", + "utoipa", + "utoipa-axum", + "utoipa-scalar", "uuid", ] @@ -3355,6 +3366,55 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" +[[package]] +name = "utoipa" +version = "5.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8bde15df68e80b16c7d16b9616e80770ad158988daa56a27dccd1e55558b0160" +dependencies = [ + "indexmap", + "serde", + "serde_json", + "utoipa-gen", +] + +[[package]] +name = "utoipa-axum" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7c25bae5bccc842449ec0c5ddc5cbb6a3a1eaeac4503895dc105a1138f8234a0" +dependencies = [ + "axum", + "paste", + "tower-layer", + "tower-service", + "utoipa", +] + +[[package]] +name = "utoipa-gen" +version = "5.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ba0b99ee52df3028635d93840c797102da61f8a7bb3cf751032455895b52ef8" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "uuid", +] + +[[package]] +name = "utoipa-scalar" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59559e1509172f6b26c1cdbc7247c4ddd1ac6560fe94b584f81ee489b141f719" +dependencies = [ + "axum", + "serde", + "serde_json", + "utoipa", +] + [[package]] name = "uuid" version = "1.23.3" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 793f6a0..93bc7e1 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -89,4 +89,7 @@ tower-sessions-redis-store = { tracing = "0.1.44" tracing-appender = "0.2.5" tracing-subscriber = { version = "0.3.23", features = ["env-filter", "json"] } +utoipa = { version = "5.5.0", features = ["chrono", "uuid"] } +utoipa-axum = "0.2.0" +utoipa-scalar = { version = "0.3.0", features = ["axum"] } uuid = { version = "1.23.3", features = ["serde", "v4"] } diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index 13f94fa..139197d 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -3,12 +3,13 @@ use axum::{ extract::{Path, Query, State}, http::StatusCode, response::{IntoResponse, Redirect}, - routing, }; use serde::Deserialize; +use utoipa_axum::{router::OpenApiRouter, routes}; use crate::{ api::RoutePrefix, + db::models::ChatRsUser, error::AppResult, extractors::{ database::Database, @@ -18,18 +19,18 @@ use crate::{ state::AppState, }; -pub fn routes() -> axum::Router { - axum::Router::new() - .route("/login/{provider}", routing::get(login_handler)) - .route("/login/{provider}/callback", routing::get(callback_handler)) - .route("/user", routing::get(get_user_handler)) - .route("/logout", routing::get(logout_handler).post(logout_handler)) +pub fn routes() -> OpenApiRouter { + OpenApiRouter::new() + .routes(routes!(login_handler, logout_handler)) + .routes(routes!(login_callback_handler)) + .routes(routes!(get_user_handler)) } fn callback_path(route_prefix: &'static str, provider: OAuthProviderEnum) -> String { format!("{route_prefix}/login/{}/callback", provider.as_str()) } +#[utoipa::path(get, path = "/login/{provider}", params(("provider" = OAuthProviderEnum, Path)), responses((status = OK)))] async fn login_handler( Path(provider): Path, Extension(RoutePrefix(prefix)): Extension, @@ -50,7 +51,8 @@ struct OAuthCallbackQuery { state: String, } -async fn callback_handler( +#[utoipa::path(get, path = "/login/{provider}/callback", params(("provider" = OAuthProviderEnum, Path)), responses((status = OK)))] +async fn login_callback_handler( Path(provider): Path, Query(query): Query, Extension(RoutePrefix(prefix)): Extension, @@ -82,6 +84,7 @@ async fn callback_handler( Ok(Redirect::to(&state.config.server.base_url)) } +#[utoipa::path(get, path = "/user", responses((status = OK, body = ChatRsUser)))] async fn get_user_handler( UserSession { user_id }: UserSession, Database(mut db): Database, @@ -91,6 +94,7 @@ async fn get_user_handler( Ok(Json(user)) } +#[utoipa::path(get, post, path = "/logout", responses((status = NO_CONTENT)))] async fn logout_handler( session: tower_sessions::Session, State(state): State, diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index b8816c6..83fc553 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -2,9 +2,10 @@ use axum::{ Json, extract::{Path, State}, response::IntoResponse, - routing, }; -use serde::Deserialize; +use serde::{Deserialize, Serialize}; +use utoipa::ToSchema; +use utoipa_axum::{router::OpenApiRouter, routes}; use uuid::Uuid; use crate::{ @@ -14,8 +15,8 @@ use crate::{ state::AppState, }; -pub fn routes() -> axum::Router { - axum::Router::new().route("/{session_id}", routing::post(chat_stream)) +pub fn routes() -> OpenApiRouter { + OpenApiRouter::new().routes(routes!(chat_stream)) } #[derive(Debug, Deserialize)] @@ -28,6 +29,7 @@ struct ChatInput { options: LlmChatOptions, } +#[utoipa::path(get, path = "/{session_id}", params(("session_id" = Uuid, Path)), responses((status = OK, body = StreamAccess)))] async fn chat_stream( UserSession { user_id }: UserSession, Path(session_id): Path, @@ -54,5 +56,14 @@ async fn chat_stream( ) .await?; - Ok(Json(stream_access)) + Ok(Json(StreamAccess { + url: stream_access.sse_url, + token: stream_access.token, + })) +} + +#[derive(Serialize, ToSchema)] +struct StreamAccess { + url: String, + token: String, } diff --git a/server-new/src/api/health.rs b/server-new/src/api/health.rs index 5998a39..8f09846 100644 --- a/server-new/src/api/health.rs +++ b/server-new/src/api/health.rs @@ -1,9 +1,12 @@ +use utoipa_axum::{router::OpenApiRouter, routes}; + use crate::state::AppState; -pub fn routes() -> axum::Router { - axum::Router::new().route("/", axum::routing::get(health_handler)) +pub fn routes() -> OpenApiRouter { + OpenApiRouter::new().routes(routes!(health_handler)) } +#[utoipa::path(get, path = "", responses((status = OK, body = &str)))] async fn health_handler() -> &'static str { "OK" } diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index a11a1f6..5fedc50 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -1,24 +1,40 @@ -use axum::{Extension, Router}; +use axum::Extension; use axum_plugin::AdHocPlugin; +use utoipa::OpenApi; +use utoipa_axum::router::OpenApiRouter; +use utoipa_scalar::{Scalar, Servable}; -use crate::state::AppState; +use crate::{services::auth::oauth::OAuthProviderEnum, state::AppState}; pub mod auth; pub mod chat; pub mod health; -/// Adds all API routes to the server under `/api` +#[derive(OpenApi)] +#[openapi( + servers((url = "/api")), + components( + schemas(OAuthProviderEnum) + ), + tags( + (name = "chat", description = "Chat routes") + ) +)] +struct ApiDoc; + +/// Adds all API routes with OpenAPI docs to the server under `/api` pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("API routes").on_setup(|router, _state| { - let api_routes = Router::new() + let (api_routes, openapi) = OpenApiRouter::with_openapi(ApiDoc::openapi()) .nest( "/auth", auth::routes().layer(Extension(RoutePrefix("/api/auth"))), ) .nest("/chat", chat::routes()) - .nest("/health", health::routes()); + .nest("/health", health::routes()) + .split_for_parts(); - Ok(router.nest("/api", api_routes)) + Ok(router.nest("/api", api_routes.merge(Scalar::with_url("/docs", openapi)))) }) } diff --git a/server-new/src/db/models/user.rs b/server-new/src/db/models/user.rs index d2bd479..fe5842e 100644 --- a/server-new/src/db/models/user.rs +++ b/server-new/src/db/models/user.rs @@ -1,12 +1,11 @@ use diesel::prelude::*; use serde::Serialize; use serde_with::skip_serializing_none; +use utoipa::ToSchema; use uuid::Uuid; -use crate::db::UtcDateTime; - #[skip_serializing_none] -#[derive(Identifiable, Queryable, Selectable, Serialize)] +#[derive(Identifiable, Queryable, Selectable, Serialize, ToSchema)] #[diesel(table_name = super::schema::users)] pub struct ChatRsUser { pub id: Uuid, @@ -18,7 +17,7 @@ pub struct ChatRsUser { pub oidc_id: Option, pub sso_username: Option, pub created_at: chrono::DateTime, - pub updated_at: UtcDateTime, + pub updated_at: chrono::DateTime, } #[derive(Insertable, Default)] diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index 4d38ef3..d9a5acf 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -7,6 +7,7 @@ use simple_oauth::{ types::{OAuthCredentials, StandardTokenResponse, UserInfo}, }; use tower_sessions::Session; +use utoipa::ToSchema; use crate::{ config::AppConfig, @@ -34,7 +35,7 @@ pub type OAuthProviderMap = HashMap>; /// Supported OAuth providers -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, ToSchema)] #[serde(rename_all = "lowercase")] pub enum OAuthProviderEnum { Github, From ed0ce5cb3467cf85642c05aa01e011eb084026da Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Thu, 2 Jul 2026 11:00:11 -0400 Subject: [PATCH 054/111] db: add openai subtype --- server-new/.gitignore | 1 + server-new/Cargo.lock | 27 ++++++- server-new/Cargo.toml | 9 ++- .../down.sql | 2 + .../up.sql | 8 ++ server-new/src/db/models/provider.rs | 54 +++---------- server-new/src/db/schema.rs | 1 + server-new/src/llm/providers/mod.rs | 2 +- server-new/src/llm/providers/openai/mod.rs | 75 +++++++------------ .../src/llm/providers/openai/response.rs | 28 +++---- server-new/src/services/provider/error.rs | 8 +- server-new/src/services/provider/mod.rs | 19 +++-- 12 files changed, 113 insertions(+), 121 deletions(-) create mode 100644 server-new/migrations/2026-07-02-133817-0000_add_openai_subtype/down.sql create mode 100644 server-new/migrations/2026-07-02-133817-0000_add_openai_subtype/up.sql diff --git a/server-new/.gitignore b/server-new/.gitignore index b4253ee..1588a67 100644 --- a/server-new/.gitignore +++ b/server-new/.gitignore @@ -1,5 +1,6 @@ # Rust /target +.diesel_lock # Env files .env* diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index e714170..da4ef72 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -2263,6 +2263,7 @@ dependencies = [ "serde_json", "serde_with", "simple-oauth", + "strum", "thiserror 2.0.18", "tinistream-client", "tokio", @@ -2607,8 +2608,9 @@ checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" [[package]] name = "simple-oauth" -version = "0.1.0" -source = "git+https://github.com/fa-sharp/simple-oauth-rs?rev=9eebc3e#9eebc3ec19dddf83643e2796b14336bb92ccedd0" +version = "0.1.0-beta" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "46d3a9a3b349db65f2bb743ed1b29ea4e866b25eee9fea0fbfecd805f676f5c0" dependencies = [ "bon", "oauth2", @@ -2681,6 +2683,27 @@ version = "0.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" +[[package]] +name = "strum" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9628de9b8791db39ceda2b119bbe13134770b56c138ec1d3af810d045c04f9bd" +dependencies = [ + "strum_macros", +] + +[[package]] +name = "strum_macros" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab85eea0270ee17587ed4156089e10b9e6880ee688791d45a905f5b1ca36f664" +dependencies = [ + "heck 0.5.0", + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "subtle" version = "2.6.1" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 93bc7e1..6355bf5 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -55,10 +55,11 @@ serde_with = { default-features = false, features = ["macros"] } -simple-oauth = { - git = "https://github.com/fa-sharp/simple-oauth-rs", - rev = "9eebc3e", - features = ["default-tls"] +simple-oauth = { version = "0.1.0-beta", features = ["default-tls"] } +strum = { + version = "0.28.0", + default-features = false, + features = ["derive", "std"] } thiserror = "2.0.18" tinistream-client = { diff --git a/server-new/migrations/2026-07-02-133817-0000_add_openai_subtype/down.sql b/server-new/migrations/2026-07-02-133817-0000_add_openai_subtype/down.sql new file mode 100644 index 0000000..3ed05a0 --- /dev/null +++ b/server-new/migrations/2026-07-02-133817-0000_add_openai_subtype/down.sql @@ -0,0 +1,2 @@ +ALTER TABLE providers +DROP COLUMN openai_subtype; diff --git a/server-new/migrations/2026-07-02-133817-0000_add_openai_subtype/up.sql b/server-new/migrations/2026-07-02-133817-0000_add_openai_subtype/up.sql new file mode 100644 index 0000000..31027f9 --- /dev/null +++ b/server-new/migrations/2026-07-02-133817-0000_add_openai_subtype/up.sql @@ -0,0 +1,8 @@ +ALTER TABLE providers +ADD COLUMN openai_subtype TEXT; + +UPDATE providers +SET + openai_subtype = 'openrouter' +WHERE + base_url = 'https://openrouter.ai/api/v1'; diff --git a/server-new/src/db/models/provider.rs b/server-new/src/db/models/provider.rs index 49411fe..16f79ad 100644 --- a/server-new/src/db/models/provider.rs +++ b/server-new/src/db/models/provider.rs @@ -1,8 +1,7 @@ -use std::str::FromStr; - use chrono::{DateTime, Utc}; use diesel::prelude::*; -use serde::{Deserialize, Serialize}; +use serde::Serialize; +use strum::{EnumString, IntoStaticStr}; use uuid::Uuid; use crate::db::models::ChatRsUser; @@ -16,7 +15,7 @@ pub struct ChatRsProvider { // #[schemars(with = "ChatRsProviderType")] pub provider_type: String, // #[schemars(with = "OpenaiSubtype")] - // pub openai_subtype: Option, + pub openai_subtype: Option, #[serde(skip)] pub user_id: Uuid, pub default_model: String, @@ -30,6 +29,7 @@ pub struct ChatRsProvider { pub struct NewChatRsProvider<'a> { pub name: &'a str, pub provider_type: &'a str, + pub openai_subtype: &'a str, pub user_id: &'a Uuid, pub base_url: Option<&'a str>, pub default_model: &'a str, @@ -46,50 +46,20 @@ pub struct UpdateChatRsProvider<'a> { } /// The API type of the provider -#[derive(Debug, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] +#[derive(Debug, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr)] +#[strum(serialize_all = "lowercase")] pub enum ChatRsProviderType { Anthropic, - Openai, + OpenAI, Ollama, Lorem, } /// The subtype for OpenAI-compatible providers -#[derive(Debug, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -pub enum OpenaiSubtype { - Openai, - Google, +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr)] +#[strum(serialize_all = "lowercase")] +pub enum OpenAISubtype { + #[default] + OpenAI, OpenRouter, - LlmGateway, -} - -#[derive(Debug, thiserror::Error)] -#[error("invalid provider type: '{0}'")] -pub struct ParseProviderTypeError(String); - -impl FromStr for ChatRsProviderType { - type Err = ParseProviderTypeError; - - fn from_str(value: &str) -> Result { - match value { - "anthropic" => Ok(ChatRsProviderType::Anthropic), - "openai" => Ok(ChatRsProviderType::Openai), - "ollama" => Ok(ChatRsProviderType::Ollama), - "lorem" => Ok(ChatRsProviderType::Lorem), - provider => Err(ParseProviderTypeError(provider.into())), - } - } -} - -impl From<&ChatRsProviderType> for &str { - fn from(value: &ChatRsProviderType) -> Self { - match value { - ChatRsProviderType::Anthropic => "anthropic", - ChatRsProviderType::Openai => "openai", - ChatRsProviderType::Ollama => "ollama", - ChatRsProviderType::Lorem => "lorem", - } - } } diff --git a/server-new/src/db/schema.rs b/server-new/src/db/schema.rs index 7b7cee5..fbae7bc 100644 --- a/server-new/src/db/schema.rs +++ b/server-new/src/db/schema.rs @@ -94,6 +94,7 @@ diesel::table! { default_model -> Text, api_key_id -> Nullable, created_at -> Timestamptz, + openai_subtype -> Nullable, } } diff --git a/server-new/src/llm/providers/mod.rs b/server-new/src/llm/providers/mod.rs index 14a6989..06f37fa 100644 --- a/server-new/src/llm/providers/mod.rs +++ b/server-new/src/llm/providers/mod.rs @@ -3,4 +3,4 @@ mod openai; mod utils; pub use lorem::LoremProvider; -pub use openai::{OpenAIProvider, OpenAIProviderConfig, OpenAIProviderFlavor}; +pub use openai::{OpenAIProvider, OpenAIProviderConfig}; diff --git a/server-new/src/llm/providers/openai/mod.rs b/server-new/src/llm/providers/openai/mod.rs index 2b86d0f..144ff8c 100644 --- a/server-new/src/llm/providers/openai/mod.rs +++ b/server-new/src/llm/providers/openai/mod.rs @@ -2,11 +2,14 @@ use futures::StreamExt; -use crate::llm::{ - error::LlmRequestError, - interface::{LlmPromptResponse, LlmProvider, LlmStreamingResponse}, - providers::utils::{self, llm_api_request}, - types::{LlmChatRequest, LlmPrompt, LlmUsage}, +use crate::{ + db::models::OpenAISubtype, + llm::{ + error::LlmRequestError, + interface::{LlmPromptResponse, LlmProvider, LlmStreamingResponse}, + providers::utils, + types::{LlmChatRequest, LlmPrompt, LlmUsage}, + }, }; mod request; @@ -17,13 +20,7 @@ use {request::*, response::*}; const OPENAI_API_BASE_URL: &str = "https://api.openai.com/v1"; const OPENROUTER_API_BASE_URL: &str = "https://openrouter.ai/api/v1"; -/// OpenAI-compatible provider behavior variants. -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum OpenAIProviderFlavor { - OpenAI, - OpenRouter, -} -impl OpenAIProviderFlavor { +impl OpenAISubtype { fn name(self) -> &'static str { match self { Self::OpenAI => "OpenAI", @@ -50,31 +47,23 @@ impl OpenAIProviderFlavor { /// Configuration for OpenAI-compatible providers. #[derive(Debug, Clone)] pub struct OpenAIProviderConfig { - flavor: OpenAIProviderFlavor, + subtype: OpenAISubtype, api_key: String, base_url: String, } impl OpenAIProviderConfig { - pub fn openai(api_key: impl Into) -> Self { - Self::new(OpenAIProviderFlavor::OpenAI, api_key, None::) - } - - pub fn openrouter(api_key: impl Into) -> Self { - Self::new(OpenAIProviderFlavor::OpenRouter, api_key, None::) - } - pub fn new( - flavor: OpenAIProviderFlavor, + subtype: OpenAISubtype, api_key: impl Into, base_url: Option>, ) -> Self { Self { - flavor, + subtype, api_key: api_key.into(), base_url: base_url .map(Into::into) - .unwrap_or_else(|| flavor.default_base_url().to_owned()) + .unwrap_or_else(|| subtype.default_base_url().to_owned()) .trim_end_matches('/') .to_owned(), } @@ -95,45 +84,37 @@ impl OpenAIProvider { config, } } - - pub fn openai(http_client: &reqwest::Client, api_key: impl Into) -> Self { - Self::new(http_client, OpenAIProviderConfig::openai(api_key)) - } - - pub fn openrouter(http_client: &reqwest::Client, api_key: impl Into) -> Self { - Self::new(http_client, OpenAIProviderConfig::openrouter(api_key)) - } } #[derive(Debug, Clone, Copy)] struct OpenAIRequestPolicy { - flavor: OpenAIProviderFlavor, + subtype: OpenAISubtype, } impl OpenAIRequestPolicy { - fn new(flavor: OpenAIProviderFlavor) -> Self { - Self { flavor } + fn new(subtype: OpenAISubtype) -> Self { + Self { subtype } } fn max_tokens(self, max_tokens: Option) -> Option { - (!self.flavor.use_max_completion_tokens()) + (!self.subtype.use_max_completion_tokens()) .then_some(max_tokens) .flatten() } fn max_completion_tokens(self, max_tokens: Option) -> Option { - self.flavor + self.subtype .use_max_completion_tokens() .then_some(max_tokens) .flatten() } fn store(self) -> Option { - self.flavor.include_store_false().then_some(false) + self.subtype.include_store_false().then_some(false) } fn stream_options(self) -> Option { - self.flavor + self.subtype .include_usage_stream_options() .then_some(OpenAIStreamOptions { include_usage: true, @@ -143,7 +124,7 @@ impl OpenAIRequestPolicy { impl LlmProvider for OpenAIProvider { fn prompt<'r>(&'r self, prompt: LlmPrompt<'r>) -> LlmPromptResponse<'r> { - let policy = OpenAIRequestPolicy::new(self.config.flavor); + let policy = OpenAIRequestPolicy::new(self.config.subtype); let request = OpenAIRequest { model: &prompt.options.model, messages: vec![OpenAIMessage { @@ -158,8 +139,8 @@ impl LlmProvider for OpenAIProvider { }; Box::pin(async move { - let provider_name = self.config.flavor.name(); - let response = llm_api_request( + let provider_name = self.config.subtype.name(); + let response = utils::llm_api_request( &self.client, provider_name, &format!("{}/chat/completions", self.config.base_url), @@ -187,7 +168,7 @@ impl LlmProvider for OpenAIProvider { } fn stream_chat<'r>(&'r self, req: LlmChatRequest<'r>) -> LlmStreamingResponse<'r> { - let policy = OpenAIRequestPolicy::new(self.config.flavor); + let policy = OpenAIRequestPolicy::new(self.config.subtype); let openai_messages = build_openai_messages(&req.messages); // let openai_tools = tools.as_ref().map(|t| build_openai_tools(t)); // @@ -204,10 +185,10 @@ impl LlmProvider for OpenAIProvider { // modalities: options.modalities.as_ref(), ..Default::default() }; - let provider_name = self.config.flavor.name(); + let provider_name = self.config.subtype.name(); Box::pin(async move { - let response = llm_api_request( + let response = utils::llm_api_request( &self.client, provider_name, &format!("{}/chat/completions", self.config.base_url), @@ -218,11 +199,11 @@ impl LlmProvider for OpenAIProvider { let stream = async_stream::stream! { let mut sse_event_stream = utils::get_sse_events(response); - let mut tool_calls: Vec = Vec::new(); + // let mut tool_calls: Vec = Vec::new(); while let Some(event) = sse_event_stream.next().await { match event { Ok(event) => { - for chunk in parse_openai_event(event, &mut tool_calls) { + for chunk in parse_openai_event(event) { yield chunk; } } diff --git a/server-new/src/llm/providers/openai/response.rs b/server-new/src/llm/providers/openai/response.rs index ab26f5c..c44deaf 100644 --- a/server-new/src/llm/providers/openai/response.rs +++ b/server-new/src/llm/providers/openai/response.rs @@ -5,7 +5,7 @@ use crate::llm::{error::LlmStreamChunkError, interface::LlmStreamChunk, types::L /// Parse chunks from an OpenAI SSE event pub fn parse_openai_event( mut event: OpenAIStreamResponse, - _tool_calls: &mut Vec, + // _tool_calls: &mut Vec, ) -> Vec> { let mut chunks = Vec::with_capacity(1); if let Some(delta) = event.choices.pop().and_then(|c| c.delta) { @@ -98,13 +98,13 @@ pub struct OpenAIResponseDelta { // pub images: Option>, } -/// OpenAI streaming tool call -#[derive(Debug, Deserialize)] -pub struct OpenAIStreamToolCall { - id: Option, - index: usize, - function: OpenAIStreamToolCallFunction, -} +// /// OpenAI streaming tool call +// #[derive(Debug, Deserialize)] +// pub struct OpenAIStreamToolCall { +// id: Option, +// index: usize, +// function: OpenAIStreamToolCallFunction, +// } // impl OpenAIStreamToolCall { // /// Convert OpenAI tool call format to ChatRsToolCall, add tool ID @@ -125,12 +125,12 @@ pub struct OpenAIStreamToolCall { // } // } -/// OpenAI streaming tool call function -#[derive(Debug, Deserialize)] -struct OpenAIStreamToolCallFunction { - name: Option, - arguments: Option, -} +// /// OpenAI streaming tool call function +// #[derive(Debug, Deserialize)] +// struct OpenAIStreamToolCallFunction { +// name: Option, +// arguments: Option, +// } // /// OpenRouter image // #[derive(Debug, Deserialize)] diff --git a/server-new/src/services/provider/error.rs b/server-new/src/services/provider/error.rs index 80202f5..fc4c271 100644 --- a/server-new/src/services/provider/error.rs +++ b/server-new/src/services/provider/error.rs @@ -1,6 +1,4 @@ -use crate::{ - db::models::ParseProviderTypeError, error::AppError, services::auth::encryption::EncryptorError, -}; +use crate::{error::AppError, services::auth::encryption::EncryptorError}; #[derive(Debug, thiserror::Error)] pub enum ProviderError { @@ -8,8 +6,8 @@ pub enum ProviderError { NotFound, #[error("missing API key")] MissingApiKey, - #[error(transparent)] - InvalidProviderType(#[from] ParseProviderTypeError), + #[error("invalid provider type: {0}")] + InvalidProviderType(#[from] strum::ParseError), #[error("error reading/writing API keys: {0}")] Encryption(#[from] EncryptorError), #[error("database error: {0}")] diff --git a/server-new/src/services/provider/mod.rs b/server-new/src/services/provider/mod.rs index a31a98d..e00f324 100644 --- a/server-new/src/services/provider/mod.rs +++ b/server-new/src/services/provider/mod.rs @@ -5,11 +5,11 @@ use uuid::Uuid; use crate::{ db::{ DbService, - models::{ChatRsProvider, ChatRsProviderType, ChatRsSecret}, + models::{ChatRsProvider, ChatRsProviderType, ChatRsSecret, OpenAISubtype}, }, llm::{ interface::LlmProvider, - providers::{LoremProvider, OpenAIProvider}, + providers::{LoremProvider, OpenAIProvider, OpenAIProviderConfig}, }, services::{auth::encryption::Encryptor, provider::error::ProviderError}, }; @@ -40,7 +40,7 @@ impl<'r> ProviderService<'r> { .find_by_id(user_id, provider_id) .await? .ok_or(ProviderError::NotFound)?; - let provider_type = ChatRsProviderType::from_str(provider.provider_type.as_str())?; + let provider_type = ChatRsProviderType::from_str(&provider.provider_type)?; Ok((provider, provider_type, api_key_secret)) } @@ -51,7 +51,7 @@ impl<'r> ProviderService<'r> { user_id: &Uuid, provider_id: i32, ) -> Result, ProviderError> { - let (_provider, provider_type, api_key_secret) = + let (provider, provider_type, api_key_secret) = self.get_provider(db, user_id, provider_id).await?; let api_key = api_key_secret .map(|secret| { @@ -62,9 +62,16 @@ impl<'r> ProviderService<'r> { let llm_provider: Arc = match provider_type { ChatRsProviderType::Lorem => Arc::new(LoremProvider::new()), - ChatRsProviderType::Openai => Arc::new(OpenAIProvider::openai( + ChatRsProviderType::OpenAI => Arc::new(OpenAIProvider::new( self.http_client, - api_key.ok_or(ProviderError::MissingApiKey)?, + OpenAIProviderConfig::new( + provider + .openai_subtype + .and_then(|s| OpenAISubtype::from_str(&s).ok()) + .unwrap_or_default(), + api_key.ok_or(ProviderError::MissingApiKey)?, + provider.base_url, + ), )), _ => todo!(), // ChatRsProviderType::Anthropic => Box::new(AnthropicProvider::new( From d50e4f3fe325ca0c5102c18f44b8540264f82a58 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Fri, 3 Jul 2026 11:29:54 -0400 Subject: [PATCH 055/111] avoid extra allocation in stream chunks --- .../src/llm/providers/openai/response.rs | 112 ++++++++++-------- 1 file changed, 61 insertions(+), 51 deletions(-) diff --git a/server-new/src/llm/providers/openai/response.rs b/server-new/src/llm/providers/openai/response.rs index c44deaf..9c95515 100644 --- a/server-new/src/llm/providers/openai/response.rs +++ b/server-new/src/llm/providers/openai/response.rs @@ -1,61 +1,71 @@ use serde::Deserialize; -use crate::llm::{error::LlmStreamChunkError, interface::LlmStreamChunk, types::LlmUsage}; +use crate::llm::{ + interface::{LlmStreamChunk, LlmStreamChunkResult}, + types::LlmUsage, +}; /// Parse chunks from an OpenAI SSE event pub fn parse_openai_event( - mut event: OpenAIStreamResponse, + event: OpenAIStreamResponse, // _tool_calls: &mut Vec, -) -> Vec> { - let mut chunks = Vec::with_capacity(1); - if let Some(delta) = event.choices.pop().and_then(|c| c.delta) { - if let Some(text) = delta.content { - chunks.push(Ok(LlmStreamChunk::Text(text))); - } - // if let Some(tool_calls_delta) = delta.tool_calls { - // for tool_call_delta in tool_calls_delta { - // if let Some(tc) = tool_calls - // .iter_mut() - // .find(|tc| tc.index == tool_call_delta.index) - // { - // if let Some(function_arguments) = tool_call_delta.function.arguments { - // *tc.function.arguments.get_or_insert_default() += &function_arguments; - // } - // if let Some(ref tool_name) = tc.function.name { - // let chunk = LlmStreamChunk::PendingToolCall(LlmPendingToolCall { - // index: tool_call_delta.index, - // tool_name: tool_name.clone(), - // }); - // chunks.push(Ok(chunk)); - // } - // } else { - // if let Some(ref tool_name) = tool_call_delta.function.name { - // let chunk = LlmStreamChunk::PendingToolCall(LlmPendingToolCall { - // index: tool_call_delta.index, - // tool_name: tool_name.clone(), - // }); - // chunks.push(Ok(chunk)); - // } - // tool_calls.push(tool_call_delta); - // } - // } - // } - // if let Some(images) = delta.images { - // chunks.push(Ok(LlmStreamChunk::Images( - // images - // .into_iter() - // .map(|image| LlmImage { - // base64_url: image.image_url.url, - // }) - // .collect(), - // ))); - // } - } - if let Some(usage) = event.usage { - chunks.push(Ok(LlmStreamChunk::Usage(usage.into()))); - } +) -> impl Iterator { + let OpenAIStreamResponse { mut choices, usage } = event; + let text = choices + .pop() + .and_then(|choice| choice.delta) + .and_then(|delta| delta.content) + .map(|text| Ok(LlmStreamChunk::Text(text))); + let usage = usage.map(|usage| Ok(LlmStreamChunk::Usage(usage.into()))); + + [text, usage].into_iter().flatten() - chunks + // if let Some(delta) = event.choices.pop().and_then(|c| c.delta) { + // if let Some(text) = delta.content { + // chunks.push(Ok(LlmStreamChunk::Text(text))); + // } + // // if let Some(tool_calls_delta) = delta.tool_calls { + // // for tool_call_delta in tool_calls_delta { + // // if let Some(tc) = tool_calls + // // .iter_mut() + // // .find(|tc| tc.index == tool_call_delta.index) + // // { + // // if let Some(function_arguments) = tool_call_delta.function.arguments { + // // *tc.function.arguments.get_or_insert_default() += &function_arguments; + // // } + // // if let Some(ref tool_name) = tc.function.name { + // // let chunk = LlmStreamChunk::PendingToolCall(LlmPendingToolCall { + // // index: tool_call_delta.index, + // // tool_name: tool_name.clone(), + // // }); + // // chunks.push(Ok(chunk)); + // // } + // // } else { + // // if let Some(ref tool_name) = tool_call_delta.function.name { + // // let chunk = LlmStreamChunk::PendingToolCall(LlmPendingToolCall { + // // index: tool_call_delta.index, + // // tool_name: tool_name.clone(), + // // }); + // // chunks.push(Ok(chunk)); + // // } + // // tool_calls.push(tool_call_delta); + // // } + // // } + // // } + // // if let Some(images) = delta.images { + // // chunks.push(Ok(LlmStreamChunk::Images( + // // images + // // .into_iter() + // // .map(|image| LlmImage { + // // base64_url: image.image_url.url, + // // }) + // // .collect(), + // // ))); + // // } + // } + // if let Some(usage) = event.usage { + // chunks.push(Ok(LlmStreamChunk::Usage(usage.into()))); + // } } /// OpenAI API response From 1bcfdfd950022b4092d7588634309ade0540de23 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Fri, 3 Jul 2026 11:48:57 -0400 Subject: [PATCH 056/111] tweaks --- server-new/config.toml | 2 ++ server-new/src/config.rs | 2 ++ server-new/src/llm/interface.rs | 4 +--- server-new/src/plugins/redis.rs | 13 ++++++------- 4 files changed, 11 insertions(+), 10 deletions(-) diff --git a/server-new/config.toml b/server-new/config.toml index 31d87ef..3c97a52 100644 --- a/server-new/config.toml +++ b/server-new/config.toml @@ -10,6 +10,8 @@ url = "postgres://localhost" [redis] url = "redis://localhost:6379" +pool_size = 4 +timeout = 10 [services] streamer_url = "http://localhost:8081" diff --git a/server-new/src/config.rs b/server-new/src/config.rs index 7b104dc..aa54d40 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -62,6 +62,8 @@ pub struct SecurityConfig { #[derive(Debug, Clone, Deserialize)] pub struct RedisConfig { pub url: String, + pub pool_size: usize, + pub timeout: u64, } /// Plugin that reads and validates configuration, and adds it to server state diff --git a/server-new/src/llm/interface.rs b/server-new/src/llm/interface.rs index f594d8d..2f9e6ac 100644 --- a/server-new/src/llm/interface.rs +++ b/server-new/src/llm/interface.rs @@ -1,10 +1,8 @@ use futures::{future::BoxFuture, stream::BoxStream}; -use crate::llm::types::LlmPrompt; - use super::{ error::{LlmRequestError, LlmStreamChunkError}, - types::{LlmChatRequest, LlmUsage}, + types::{LlmChatRequest, LlmPrompt, LlmUsage}, }; /// Trait representing an LLM provider diff --git a/server-new/src/plugins/redis.rs b/server-new/src/plugins/redis.rs index ec334b5..8c15de7 100644 --- a/server-new/src/plugins/redis.rs +++ b/server-new/src/plugins/redis.rs @@ -6,24 +6,23 @@ use fred::prelude::*; use crate::{config::AppConfig, state::AppState}; -const DEFAULT_TIMEOUT: Duration = Duration::from_secs(8); -const POOL_SIZE: usize = 4; - pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("Redis") .on_init(async |mut state| { let app_config = state.get::().context("no config")?; let config = Config::from_url(&app_config.redis.url).context("invalid Redis URL")?; + let timeout = Duration::from_secs(app_config.redis.timeout); + let pool = Builder::from_config(config) .with_connection_config(|c| { - c.connection_timeout = DEFAULT_TIMEOUT; - c.internal_command_timeout = DEFAULT_TIMEOUT; + c.connection_timeout = timeout; + c.internal_command_timeout = timeout; c.tcp.nodelay = Some(true); }) .with_performance_config(|c| { - c.default_command_timeout = DEFAULT_TIMEOUT; + c.default_command_timeout = timeout; }) - .build_pool(POOL_SIZE)?; + .build_pool(app_config.redis.pool_size)?; pool.init().await.context("failed to connect to Redis")?; tracing::info!("Connected to Redis"); From 1160e04739e5cf195e94bcb3202c655fd77dfb68 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Fri, 3 Jul 2026 14:17:47 -0400 Subject: [PATCH 057/111] add auth config --- server-new/src/api/auth.rs | 29 +++++++----- server-new/src/extractors/auth_config.rs | 52 ++++++++++++++++++++++ server-new/src/extractors/mod.rs | 1 + server-new/src/services/auth/oauth/oidc.rs | 2 +- 4 files changed, 72 insertions(+), 12 deletions(-) create mode 100644 server-new/src/extractors/auth_config.rs diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index 139197d..53a4e80 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -12,6 +12,7 @@ use crate::{ db::models::ChatRsUser, error::AppResult, extractors::{ + auth_config::PublicAuthConfig, database::Database, session::{SessionMeta, UserSession}, }, @@ -21,9 +22,25 @@ use crate::{ pub fn routes() -> OpenApiRouter { OpenApiRouter::new() + .routes(routes!(get_user_handler)) + .routes(routes!(get_config_handler)) .routes(routes!(login_handler, logout_handler)) .routes(routes!(login_callback_handler)) - .routes(routes!(get_user_handler)) +} + +#[utoipa::path(get, path = "/user", responses((status = OK, body = ChatRsUser)))] +async fn get_user_handler( + UserSession { user_id }: UserSession, + Database(mut db): Database, + State(state): State, +) -> AppResult { + let user = state.auth_service().get_user(&mut db, &user_id).await?; + Ok(Json(user)) +} + +#[utoipa::path(get, path = "/config", responses((status = OK, body = PublicAuthConfig)))] +async fn get_config_handler(auth_config: PublicAuthConfig) -> impl IntoResponse { + Json(auth_config) } fn callback_path(route_prefix: &'static str, provider: OAuthProviderEnum) -> String { @@ -84,16 +101,6 @@ async fn login_callback_handler( Ok(Redirect::to(&state.config.server.base_url)) } -#[utoipa::path(get, path = "/user", responses((status = OK, body = ChatRsUser)))] -async fn get_user_handler( - UserSession { user_id }: UserSession, - Database(mut db): Database, - State(state): State, -) -> AppResult { - let user = state.auth_service().get_user(&mut db, &user_id).await?; - Ok(Json(user)) -} - #[utoipa::path(get, post, path = "/logout", responses((status = NO_CONTENT)))] async fn logout_handler( session: tower_sessions::Session, diff --git a/server-new/src/extractors/auth_config.rs b/server-new/src/extractors/auth_config.rs new file mode 100644 index 0000000..3dfd0a3 --- /dev/null +++ b/server-new/src/extractors/auth_config.rs @@ -0,0 +1,52 @@ +use axum::extract::FromRequestParts; +use serde::Serialize; +use utoipa::ToSchema; + +use crate::state::AppState; + +/// The current auth configuration of the server +#[derive(Debug, Serialize, ToSchema)] +pub struct PublicAuthConfig { + /// Whether GitHub login is enabled + github: bool, + /// Whether Google login is enabled + google: bool, + /// Whether Discord login is enabled + discord: bool, + /// OIDC configuration + oidc: Option, + // /// SSO configuration + // sso: Option, +} + +#[derive(Debug, Serialize, ToSchema)] +struct Oidc { + /// The name of the OIDC provider + name: String, +} + +// #[derive(Debug, JsonSchema, serde::Serialize)] +// struct SSO { +// /// Whether SSO header authentication is enabled +// enabled: bool, +// /// The URL to redirect to after logout +// logout_url: Option, +// } + +impl FromRequestParts for PublicAuthConfig { + type Rejection = (); + + async fn from_request_parts( + _parts: &mut axum::http::request::Parts, + state: &AppState, + ) -> Result { + Ok(PublicAuthConfig { + github: state.config.auth.github.is_some(), + discord: state.config.auth.discord.is_some(), + google: state.config.auth.google.is_some(), + oidc: state.config.auth.oidc.as_ref().map(|oidc| Oidc { + name: oidc.name.as_deref().unwrap_or("OIDC").to_owned(), + }), + }) + } +} diff --git a/server-new/src/extractors/mod.rs b/server-new/src/extractors/mod.rs index 26d219c..5988553 100644 --- a/server-new/src/extractors/mod.rs +++ b/server-new/src/extractors/mod.rs @@ -1,2 +1,3 @@ +pub mod auth_config; pub mod database; pub mod session; diff --git a/server-new/src/services/auth/oauth/oidc.rs b/server-new/src/services/auth/oauth/oidc.rs index a668b66..6d4eef6 100644 --- a/server-new/src/services/auth/oauth/oidc.rs +++ b/server-new/src/services/auth/oauth/oidc.rs @@ -14,7 +14,7 @@ use super::OAuthProvider; #[derive(Clone, Debug, Deserialize)] pub struct OidcConfig { - name: Option, + pub name: Option, client_id: String, client_secret: String, auth_endpoint: String, From 55cbe403b8909694fdd71d9ff6032329a60f705b Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sat, 4 Jul 2026 02:22:12 -0400 Subject: [PATCH 058/111] add api keys --- server-new/src/api/api_key.rs | 85 +++++++++++++++++++++++ server-new/src/api/auth.rs | 81 ++++++++++++++------- server-new/src/api/chat.rs | 15 +++- server-new/src/api/mod.rs | 17 ++++- server-new/src/db/mod.rs | 11 +-- server-new/src/db/models.rs | 4 +- server-new/src/db/models/api_key.rs | 25 +++++++ server-new/src/db/repositories.rs | 2 + server-new/src/db/repositories/api_key.rs | 70 +++++++++++++++++++ server-new/src/error.rs | 6 ++ server-new/src/extractors/auth_config.rs | 2 + server-new/src/extractors/mod.rs | 14 +++- server-new/src/extractors/session.rs | 61 +--------------- server-new/src/extractors/user.rs | 84 ++++++++++++++++++++++ server-new/src/services/auth/api_key.rs | 80 +++++++++++++++++++++ server-new/src/services/auth/error.rs | 4 +- server-new/src/services/auth/mod.rs | 25 +++++-- server-new/src/services/auth/oauth.rs | 20 ++---- server-new/src/services/auth/session.rs | 10 +-- server-new/src/state.rs | 2 +- 20 files changed, 491 insertions(+), 127 deletions(-) create mode 100644 server-new/src/api/api_key.rs create mode 100644 server-new/src/db/models/api_key.rs create mode 100644 server-new/src/db/repositories/api_key.rs create mode 100644 server-new/src/extractors/user.rs create mode 100644 server-new/src/services/auth/api_key.rs diff --git a/server-new/src/api/api_key.rs b/server-new/src/api/api_key.rs new file mode 100644 index 0000000..6417d97 --- /dev/null +++ b/server-new/src/api/api_key.rs @@ -0,0 +1,85 @@ +use axum::{ + Json, + extract::{Path, State}, + http::StatusCode, +}; +use serde::{Deserialize, Serialize}; +use utoipa::ToSchema; +use utoipa_axum::{router::OpenApiRouter, routes}; +use uuid::Uuid; + +use crate::{ + api::ApiTag, + db::models::ChatRsApiKey, + error::AppResult, + extractors::{CurrentUser, Database}, + state::AppState, +}; + +pub fn routes() -> OpenApiRouter { + OpenApiRouter::new().routes(routes!(list_api_keys, create_api_key, delete_api_key)) +} + +/// List all API keys +#[utoipa::path( + get, path = "", + responses((status = OK, body = Vec)), + tag = ApiTag::ApiKey.into()) +] +async fn list_api_keys( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, +) -> AppResult>> { + let keys = db.api_keys().find_by_user_id(&user_id).await?; + Ok(Json(keys)) +} + +#[derive(Deserialize, ToSchema)] +struct ApiKeyCreateInput { + name: String, +} + +#[derive(Serialize, ToSchema)] +struct ApiKeyCreateResponse { + id: Uuid, + key: String, +} + +/// Create an API key +#[utoipa::path( + post, path = "", + request_body = ApiKeyCreateInput, + responses((status = OK, body = ApiKeyCreateResponse)), + tag = ApiTag::ApiKey.into()) +] +async fn create_api_key( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, + State(state): State, + input: Json, +) -> AppResult> { + let (id, key) = state + .auth_service() + .api_keys() + .create_api_key(&mut db, &user_id, &input.name) + .await?; + + Ok(Json(ApiKeyCreateResponse { id, key })) +} + +/// Delete an API key +#[utoipa::path( + delete, path = "/{id}", + params(("id" = Uuid, Path)), + responses((status = NO_CONTENT)), + tag = ApiTag::ApiKey.into()) +] +async fn delete_api_key( + CurrentUser { user_id }: CurrentUser, + Path(api_key_id): Path, + Database(mut db): Database, +) -> AppResult { + let _deleted_id = db.api_keys().delete(&user_id, &api_key_id).await?; + + Ok(StatusCode::NO_CONTENT) +} diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index 53a4e80..9652f24 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -5,32 +5,35 @@ use axum::{ response::{IntoResponse, Redirect}, }; use serde::Deserialize; +use utoipa::{IntoParams, ToSchema}; use utoipa_axum::{router::OpenApiRouter, routes}; use crate::{ - api::RoutePrefix, + api::{ApiTag, RoutePrefix}, db::models::ChatRsUser, error::AppResult, - extractors::{ - auth_config::PublicAuthConfig, - database::Database, - session::{SessionMeta, UserSession}, - }, + extractors::{CurrentUser, Database, PublicAuthConfig, SessionMeta}, services::auth::oauth::OAuthProviderEnum, state::AppState, }; pub fn routes() -> OpenApiRouter { OpenApiRouter::new() - .routes(routes!(get_user_handler)) - .routes(routes!(get_config_handler)) - .routes(routes!(login_handler, logout_handler)) - .routes(routes!(login_callback_handler)) + .routes(routes!(get_user)) + .routes(routes!(get_config)) + .routes(routes!(oauth_login)) + .routes(routes!(oauth_login_callback)) + .routes(routes!(logout)) } -#[utoipa::path(get, path = "/user", responses((status = OK, body = ChatRsUser)))] -async fn get_user_handler( - UserSession { user_id }: UserSession, +/// Get current user +#[utoipa::path( + get, path = "/user", + responses((status = OK, body = ChatRsUser)), + tag = ApiTag::Auth.into()) +] +async fn get_user( + CurrentUser { user_id }: CurrentUser, Database(mut db): Database, State(state): State, ) -> AppResult { @@ -38,17 +41,27 @@ async fn get_user_handler( Ok(Json(user)) } -#[utoipa::path(get, path = "/config", responses((status = OK, body = PublicAuthConfig)))] -async fn get_config_handler(auth_config: PublicAuthConfig) -> impl IntoResponse { +#[utoipa::path( + get, path = "/config", + responses((status = OK, body = PublicAuthConfig)), + tag = ApiTag::Auth.into() +)] +async fn get_config(auth_config: PublicAuthConfig) -> impl IntoResponse { Json(auth_config) } -fn callback_path(route_prefix: &'static str, provider: OAuthProviderEnum) -> String { - format!("{route_prefix}/login/{}/callback", provider.as_str()) +fn oauth_callback_path(route_prefix: &'static str, provider: OAuthProviderEnum) -> String { + format!("{route_prefix}/login/{provider}/callback") } -#[utoipa::path(get, path = "/login/{provider}", params(("provider" = OAuthProviderEnum, Path)), responses((status = OK)))] -async fn login_handler( +/// OAuth login redirect +#[utoipa::path( + get, path = "/login/{provider}", + params(("provider" = OAuthProviderEnum, Path)), + responses((status = OK)), + tag = ApiTag::Auth.into(), +)] +async fn oauth_login( Path(provider): Path, Extension(RoutePrefix(prefix)): Extension, State(state): State, @@ -56,20 +69,29 @@ async fn login_handler( ) -> AppResult { let oauth = state.auth_service().oauth(); let auth_url = oauth - .authorize_url(provider, &callback_path(prefix, provider), &session) + .authorize_url(provider, &oauth_callback_path(prefix, provider), &session) .await?; Ok(Redirect::to(auth_url.as_str())) } -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Deserialize, IntoParams, ToSchema)] struct OAuthCallbackQuery { code: String, state: String, } -#[utoipa::path(get, path = "/login/{provider}/callback", params(("provider" = OAuthProviderEnum, Path)), responses((status = OK)))] -async fn login_callback_handler( +/// OAuth login callback +#[utoipa::path( + get, path = "/login/{provider}/callback", + params( + ("query" = inline(OAuthCallbackQuery), Query), + ("provider" = OAuthProviderEnum, Path, description = "the OAuth provider") + ), + responses((status = OK)), + tag = ApiTag::Auth.into(), +)] +async fn oauth_login_callback( Path(provider): Path, Query(query): Query, Extension(RoutePrefix(prefix)): Extension, @@ -77,13 +99,13 @@ async fn login_callback_handler( State(state): State, session: tower_sessions::Session, meta: SessionMeta, - maybe_user: Option, + maybe_user: Option, ) -> AppResult { let oauth = state.auth_service().oauth(); let token = oauth .exchange_code( provider, - &callback_path(prefix, provider), + &oauth_callback_path(prefix, provider), &session, &query.code, &query.state, @@ -101,8 +123,13 @@ async fn login_callback_handler( Ok(Redirect::to(&state.config.server.base_url)) } -#[utoipa::path(get, post, path = "/logout", responses((status = NO_CONTENT)))] -async fn logout_handler( +/// Logout +#[utoipa::path( + method(get, post), path = "/logout", + tag = ApiTag::Auth.into(), + responses((status = NO_CONTENT)), +)] +async fn logout( session: tower_sessions::Session, State(state): State, ) -> AppResult { diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index 83fc553..847cdf1 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -9,8 +9,9 @@ use utoipa_axum::{router::OpenApiRouter, routes}; use uuid::Uuid; use crate::{ + api::ApiTag, error::AppError, - extractors::{database::Database, session::UserSession}, + extractors::{CurrentUser, Database}, llm::types::{LlmChatOptions, LlmUserMessage}, state::AppState, }; @@ -29,9 +30,17 @@ struct ChatInput { options: LlmChatOptions, } -#[utoipa::path(get, path = "/{session_id}", params(("session_id" = Uuid, Path)), responses((status = OK, body = StreamAccess)))] +/// Streaming chat +/// +/// Send a message in a chat session and stream the response +#[utoipa::path( + get, path = "/{session_id}", + params(("session_id" = Uuid, Path)), + responses((status = OK, body = StreamAccess)), + tag = ApiTag::Chat.into(), +)] async fn chat_stream( - UserSession { user_id }: UserSession, + CurrentUser { user_id }: CurrentUser, Path(session_id): Path, Database(mut db): Database, State(state): State, diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index 5fedc50..ccb1bfb 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -1,15 +1,25 @@ use axum::Extension; use axum_plugin::AdHocPlugin; +use strum::{AsRefStr, IntoStaticStr}; use utoipa::OpenApi; use utoipa_axum::router::OpenApiRouter; use utoipa_scalar::{Scalar, Servable}; use crate::{services::auth::oauth::OAuthProviderEnum, state::AppState}; +pub mod api_key; pub mod auth; pub mod chat; pub mod health; +#[derive(AsRefStr, IntoStaticStr)] +#[strum(serialize_all = "snake_case")] +enum ApiTag { + ApiKey, + Auth, + Chat, +} + #[derive(OpenApi)] #[openapi( servers((url = "/api")), @@ -17,7 +27,9 @@ pub mod health; schemas(OAuthProviderEnum) ), tags( - (name = "chat", description = "Chat routes") + (name = ApiTag::ApiKey.as_ref(), description = "Manage API keys"), + (name = ApiTag::Auth.as_ref(), description = "Authentication"), + (name = ApiTag::Chat.as_ref(), description = "Chats and sessions") ) )] struct ApiDoc; @@ -26,6 +38,7 @@ struct ApiDoc; pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("API routes").on_setup(|router, _state| { let (api_routes, openapi) = OpenApiRouter::with_openapi(ApiDoc::openapi()) + .nest("/api_key", api_key::routes()) .nest( "/auth", auth::routes().layer(Extension(RoutePrefix("/api/auth"))), @@ -39,4 +52,4 @@ pub fn plugin() -> AdHocPlugin { } #[derive(Clone)] -struct RoutePrefix(pub &'static str); +struct RoutePrefix(&'static str); diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs index cab508e..c1174fc 100644 --- a/server-new/src/db/mod.rs +++ b/server-new/src/db/mod.rs @@ -15,8 +15,8 @@ pub type DbPool = Pool; /// Error when attempting to retrieve a connection from the pool pub type DbPoolError = PoolError; -/// The database connection retrieved from the pool. -/// For pipelining multiple queries, can use a shared reference with `&mut &**conn`. +/// The database connection retrieved from the pool. For pipelining multiple +/// queries in Diesel, a shared reference can be used with `&mut &**conn`. pub struct DbConnection(Object); impl Deref for DbConnection { type Target = AsyncPgConnection; @@ -48,8 +48,8 @@ impl DbService { Ok(Self::new(DbConnection(cxn))) } - pub fn users(&mut self) -> repositories::UserRepository<'_> { - repositories::UserRepository::new(&mut self.cxn) + pub fn api_keys(&mut self) -> repositories::ApiKeyRepository<'_> { + repositories::ApiKeyRepository::new(&mut self.cxn) } pub fn auth_sessions(&mut self) -> repositories::SessionRepository<'_> { repositories::SessionRepository::new(&mut self.cxn) @@ -60,4 +60,7 @@ impl DbService { pub fn providers(&mut self) -> repositories::ProviderRepository<'_> { repositories::ProviderRepository::new(&mut self.cxn) } + pub fn users(&mut self) -> repositories::UserRepository<'_> { + repositories::UserRepository::new(&mut self.cxn) + } } diff --git a/server-new/src/db/models.rs b/server-new/src/db/models.rs index feae9a2..2b12e99 100644 --- a/server-new/src/db/models.rs +++ b/server-new/src/db/models.rs @@ -1,6 +1,6 @@ use crate::db::schema; -// mod api_key; +mod api_key; mod chat; // mod file; mod provider; @@ -9,7 +9,7 @@ mod secret; mod session; mod user; -// pub use api_key::*; +pub use api_key::*; pub use chat::*; // pub use file::*; pub use provider::*; diff --git a/server-new/src/db/models/api_key.rs b/server-new/src/db/models/api_key.rs new file mode 100644 index 0000000..2cc1363 --- /dev/null +++ b/server-new/src/db/models/api_key.rs @@ -0,0 +1,25 @@ +use chrono::{DateTime, Utc}; +use diesel::prelude::*; +use serde::Serialize; +use utoipa::ToSchema; +use uuid::Uuid; + +use crate::db::models::ChatRsUser; + +#[derive(Identifiable, Queryable, Selectable, Associations, Serialize, ToSchema)] +#[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] +#[diesel(table_name = super::schema::app_api_keys)] +pub struct ChatRsApiKey { + pub id: Uuid, + #[serde(skip)] + pub user_id: Uuid, + pub name: String, + pub created_at: DateTime, +} + +#[derive(Insertable)] +#[diesel(table_name = super::schema::app_api_keys)] +pub struct NewChatRsApiKey<'r> { + pub user_id: &'r Uuid, + pub name: &'r str, +} diff --git a/server-new/src/db/repositories.rs b/server-new/src/db/repositories.rs index 16624ea..c5b4d12 100644 --- a/server-new/src/db/repositories.rs +++ b/server-new/src/db/repositories.rs @@ -1,8 +1,10 @@ +mod api_key; mod chat; mod provider; mod session; mod user; +pub use api_key::ApiKeyRepository; pub use chat::ChatRepository; pub use provider::ProviderRepository; pub use session::SessionRepository; diff --git a/server-new/src/db/repositories/api_key.rs b/server-new/src/db/repositories/api_key.rs new file mode 100644 index 0000000..42aed9a --- /dev/null +++ b/server-new/src/db/repositories/api_key.rs @@ -0,0 +1,70 @@ +use diesel::prelude::*; +use diesel::result::Error; +use diesel_async::RunQueryDsl; +use uuid::Uuid; + +use crate::db::{ + DbConnection, + models::{ChatRsApiKey, NewChatRsApiKey}, + schema::app_api_keys, +}; + +pub struct ApiKeyRepository<'a> { + pub db: &'a mut DbConnection, +} + +impl<'a> ApiKeyRepository<'a> { + pub fn new(db: &'a mut DbConnection) -> Self { + ApiKeyRepository { db } + } + + pub async fn find_by_id(&mut self, id: &Uuid) -> Result, Error> { + app_api_keys::table + .find(id) + .select(ChatRsApiKey::as_select()) + .first(self.db) + .await + .optional() + } + + pub async fn find_by_user_id(&mut self, user_id: &Uuid) -> Result, Error> { + let keys = app_api_keys::table + .filter(app_api_keys::user_id.eq(user_id)) + .select(ChatRsApiKey::as_select()) + .load(self.db) + .await?; + + Ok(keys) + } + + pub async fn create(&mut self, api_key: NewChatRsApiKey<'_>) -> Result { + let id: Uuid = diesel::insert_into(app_api_keys::table) + .values(api_key) + .returning(app_api_keys::id) + .get_result(self.db) + .await?; + + Ok(id) + } + + pub async fn delete(&mut self, user_id: &Uuid, api_key_id: &Uuid) -> Result { + let id: Uuid = diesel::delete(app_api_keys::table) + .filter(app_api_keys::id.eq(api_key_id)) + .filter(app_api_keys::user_id.eq(user_id)) + .returning(app_api_keys::id) + .get_result(self.db) + .await?; + + Ok(id) + } + + pub async fn delete_by_user(&mut self, user_id: &Uuid) -> Result, Error> { + let ids: Vec = diesel::delete(app_api_keys::table) + .filter(app_api_keys::user_id.eq(user_id)) + .returning(app_api_keys::id) + .get_results(self.db) + .await?; + + Ok(ids) + } +} diff --git a/server-new/src/error.rs b/server-new/src/error.rs index dba225d..3198802 100644 --- a/server-new/src/error.rs +++ b/server-new/src/error.rs @@ -50,6 +50,12 @@ impl AppError { } } +impl From for AppError { + fn from(err: diesel::result::Error) -> Self { + Self::internal(anyhow::Error::from(err).context("database error")) + } +} + #[derive(Debug, Serialize)] struct ErrorResponse { error: ErrorBody, diff --git a/server-new/src/extractors/auth_config.rs b/server-new/src/extractors/auth_config.rs index 3dfd0a3..f0adb49 100644 --- a/server-new/src/extractors/auth_config.rs +++ b/server-new/src/extractors/auth_config.rs @@ -1,10 +1,12 @@ use axum::extract::FromRequestParts; use serde::Serialize; +use serde_with::skip_serializing_none; use utoipa::ToSchema; use crate::state::AppState; /// The current auth configuration of the server +#[skip_serializing_none] #[derive(Debug, Serialize, ToSchema)] pub struct PublicAuthConfig { /// Whether GitHub login is enabled diff --git a/server-new/src/extractors/mod.rs b/server-new/src/extractors/mod.rs index 5988553..1ffb082 100644 --- a/server-new/src/extractors/mod.rs +++ b/server-new/src/extractors/mod.rs @@ -1,3 +1,11 @@ -pub mod auth_config; -pub mod database; -pub mod session; +//! Extractors to be used in API route handlers + +mod auth_config; +mod database; +mod session; +mod user; + +pub use auth_config::PublicAuthConfig; +pub use database::Database; +pub use session::SessionMeta; +pub use user::CurrentUser; diff --git a/server-new/src/extractors/session.rs b/server-new/src/extractors/session.rs index c902367..368207a 100644 --- a/server-new/src/extractors/session.rs +++ b/server-new/src/extractors/session.rs @@ -4,33 +4,12 @@ use std::{ }; use axum::{ - extract::{ConnectInfo, FromRequestParts, OptionalFromRequestParts}, + extract::{ConnectInfo, FromRequestParts}, http::header, }; -use tower_sessions::Session; - -use crate::{error::AppError, state::AppState}; - use serde::{Deserialize, Serialize}; -use uuid::Uuid; - -use crate::db::UtcDateTime; - -/// Represents an active user session. This can be used as an extractor in route handlers: -/// - If used as `UserSession`, request will automatically return an unauthorized error -/// if there is no active session. -/// - If used as `Option`, will be `Some` if there is an active session -/// and `None` otherwise. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct UserSession { - pub user_id: Uuid, -} -impl UserSession { - pub fn new(user_id: Uuid) -> Self { - Self { user_id } - } -} +use crate::{db::UtcDateTime, state::AppState}; /// Session metadata extracted on login. #[derive(Debug, Clone, Serialize, Deserialize)] @@ -40,42 +19,6 @@ pub struct SessionMeta { pub user_agent: Option, } -/// Active user session data. -impl OptionalFromRequestParts for UserSession { - type Rejection = AppError; - - async fn from_request_parts( - parts: &mut axum::http::request::Parts, - state: &AppState, - ) -> Result, Self::Rejection> { - let session = Session::from_request_parts(parts, state) - .await - .map_err(|(_, msg)| AppError::internal(anyhow::anyhow!(msg)))?; - let user_session = state - .auth_service() - .session() - .user_session(&session) - .await?; - - Ok(user_session) - } -} - -impl FromRequestParts for UserSession { - type Rejection = AppError; - - async fn from_request_parts( - parts: &mut axum::http::request::Parts, - state: &AppState, - ) -> Result { - match >::from_request_parts(parts, state).await? - { - Some(user_session) => Ok(user_session), - None => Err(AppError::unauthorized("no active session")), - } - } -} - impl FromRequestParts for SessionMeta { type Rejection = (); diff --git a/server-new/src/extractors/user.rs b/server-new/src/extractors/user.rs new file mode 100644 index 0000000..c7200e9 --- /dev/null +++ b/server-new/src/extractors/user.rs @@ -0,0 +1,84 @@ +use anyhow::anyhow; +use axum::{ + extract::{FromRequestParts, OptionalFromRequestParts}, + http::header, +}; +use serde::{Deserialize, Serialize}; +use uuid::Uuid; + +use crate::{db::DbService, error::AppError, state::AppState}; + +/** +Represents an active user, extracted from the session or API key. This can be used +as an extractor in route handlers: +- If used as `CurrentUser`, request will automatically return an unauthorized error +if there is no active user. +- If used as `Option`, will be `Some` if there is an active user +and `None` otherwise. +*/ +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct CurrentUser { + pub user_id: Uuid, +} + +impl CurrentUser { + fn new(user_id: Uuid) -> Self { + Self { user_id } + } +} + +impl OptionalFromRequestParts for CurrentUser { + type Rejection = AppError; + + async fn from_request_parts( + parts: &mut axum::http::request::Parts, + state: &AppState, + ) -> Result, Self::Rejection> { + // If there is an Authorization header, validate the API key + if let Some(auth_header) = parts + .headers + .get(header::AUTHORIZATION) + .and_then(|bytes| bytes.to_str().ok()) + { + let mut db = DbService::from_pool(&state.db_pool) + .await + .map_err(|err| AppError::internal(err.into()))?; + let user_id = state + .auth_service() + .api_keys() + .validate_api_key(&mut db, auth_header) + .await?; + + Ok(Some(Self { user_id })) + } + // Check for session + else { + let session = parts + .extensions + .get::() + .ok_or_else(|| AppError::internal(anyhow!("session not attached to request")))?; + let maybe_user_id = state + .auth_service() + .session() + .active_user_id(&session) + .await?; + + Ok(maybe_user_id.map(Self::new)) + } + } +} + +impl FromRequestParts for CurrentUser { + type Rejection = AppError; + + async fn from_request_parts( + parts: &mut axum::http::request::Parts, + state: &AppState, + ) -> Result { + match >::from_request_parts(parts, state).await? + { + Some(user_session) => Ok(user_session), + None => Err(AppError::unauthorized("no active session")), + } + } +} diff --git a/server-new/src/services/auth/api_key.rs b/server-new/src/services/auth/api_key.rs new file mode 100644 index 0000000..ae79d90 --- /dev/null +++ b/server-new/src/services/auth/api_key.rs @@ -0,0 +1,80 @@ +use uuid::Uuid; + +use crate::{ + db::{DbService, models::NewChatRsApiKey}, + services::auth::{ + encryption::Encryptor, + error::{AuthError, AuthResult}, + }, +}; + +const API_KEY_PREFIX: &str = "rs-chat-key"; +const API_KEY_HEADER_PREFIX: &str = "Bearer rs-chat-key|"; + +pub struct ApiKeyService<'r> { + encryptor: &'r Encryptor, +} + +impl<'r> ApiKeyService<'r> { + pub fn new(encryptor: &'r Encryptor) -> Self { + Self { encryptor } + } + + /// Build an API key string from the given ciphertext and nonce + fn build_api_key(ciphertext: &[u8], nonce: &[u8]) -> String { + format!( + "{API_KEY_PREFIX}|{}|{}", + hex::encode(nonce), + hex::encode(ciphertext) + ) + } + + /// Create an API key, and return its ID and the encrypted key + pub async fn create_api_key( + &self, + db: &mut DbService, + user_id: &Uuid, + name: &str, + ) -> AuthResult<(Uuid, String)> { + let key_id = db + .api_keys() + .create(NewChatRsApiKey { + user_id: &user_id, + name: &name, + }) + .await?; + let (ciphertext, nonce) = self.encryptor.encrypt_bytes(key_id.as_bytes())?; + + Ok((key_id, Self::build_api_key(&ciphertext, &nonce))) + } + + /// Validate the API key and get the user ID + pub async fn validate_api_key( + &self, + db: &mut DbService, + auth_header: &str, + ) -> AuthResult { + let (nonce, ciphertext) = auth_header + .strip_prefix(API_KEY_HEADER_PREFIX) + .and_then(|s| s.split_once('|')) + .and_then(|(nonce_hex, cipher_hex)| { + hex::decode(nonce_hex) + .ok() + .zip(hex::decode(cipher_hex).ok()) + }) + .ok_or(AuthError::Unauthorized("invalid API key format"))?; + let api_key_id = self + .encryptor + .decrypt_bytes(&ciphertext, &nonce) + .map_err(|_| AuthError::Unauthorized("failed to decrypt API key")) + .and_then(|key_bytes| { + Uuid::from_slice(&key_bytes) + .map_err(|_| AuthError::Unauthorized("couldn't parse API key id")) + })?; + + match db.api_keys().find_by_id(&api_key_id).await? { + Some(api_key) => Ok(api_key.user_id), + None => Err(AuthError::Unauthorized("API key not found")), + } + } +} diff --git a/server-new/src/services/auth/error.rs b/server-new/src/services/auth/error.rs index aa9ca27..9060fdd 100644 --- a/server-new/src/services/auth/error.rs +++ b/server-new/src/services/auth/error.rs @@ -1,4 +1,4 @@ -use crate::error::AppError; +use crate::{error::AppError, services::auth::encryption::EncryptorError}; pub type AuthResult = Result; @@ -13,6 +13,8 @@ pub enum AuthError { OAuth(#[from] simple_oauth::SimpleOAuthError), #[error("user not found")] UserNotFound, + #[error("encryption error: {0}")] + Encryption(#[from] EncryptorError), #[error("database error: {0}")] Database(#[from] diesel::result::Error), #[error("database pool error: {0}")] diff --git a/server-new/src/services/auth/mod.rs b/server-new/src/services/auth/mod.rs index b071a10..bda9f1a 100644 --- a/server-new/src/services/auth/mod.rs +++ b/server-new/src/services/auth/mod.rs @@ -1,9 +1,11 @@ use crate::{ config::AppConfig, db::{DbService, models::ChatRsUser}, + services::auth::encryption::Encryptor, }; use uuid::Uuid; +pub mod api_key; pub mod encryption; mod error; pub mod oauth; @@ -12,15 +14,21 @@ pub mod session_store; use error::{AuthError, AuthResult}; -pub struct AuthService<'a> { - config: &'a AppConfig, - oauth_providers: &'a oauth::OAuthProviderMap, +pub struct AuthService<'r> { + config: &'r AppConfig, + encryptor: &'r Encryptor, + oauth_providers: &'r oauth::OAuthProviderMap, } -impl<'a> AuthService<'a> { - pub fn new(config: &'a AppConfig, oauth_providers: &'a oauth::OAuthProviderMap) -> Self { +impl<'r> AuthService<'r> { + pub fn new( + config: &'r AppConfig, + encryptor: &'r Encryptor, + oauth_providers: &'r oauth::OAuthProviderMap, + ) -> Self { Self { config, + encryptor, oauth_providers, } } @@ -40,7 +48,12 @@ impl<'a> AuthService<'a> { } /// Access OAuth functions - pub fn oauth(self) -> oauth::OAuthService<'a> { + pub fn oauth(self) -> oauth::OAuthService<'r> { oauth::OAuthService::new(self.config, self.oauth_providers) } + + /// Access API key functions + pub fn api_keys(self) -> api_key::ApiKeyService<'r> { + api_key::ApiKeyService::new(self.encryptor) + } } diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index d9a5acf..fe0fdc9 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -6,6 +6,7 @@ use simple_oauth::{ SimpleOAuthClient, SimpleOAuthError, SimpleOAuthProvider, types::{OAuthCredentials, StandardTokenResponse, UserInfo}, }; +use strum::Display; use tower_sessions::Session; use utoipa::ToSchema; @@ -15,7 +16,7 @@ use crate::{ DbService, models::{ChatRsUser, NewChatRsUser, UpdateChatRsUser}, }, - extractors::session::UserSession, + extractors::CurrentUser, services::auth::{AuthError, AuthResult}, }; @@ -34,25 +35,16 @@ pub type OAuthProviderMap = HashMap>; -/// Supported OAuth providers -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, ToSchema)] +/// Supported OAuth provider +#[derive(Debug, Display, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, ToSchema)] #[serde(rename_all = "lowercase")] +#[strum(serialize_all = "lowercase")] pub enum OAuthProviderEnum { Github, Discord, Google, Oidc, } -impl OAuthProviderEnum { - pub fn as_str(&self) -> &str { - match self { - OAuthProviderEnum::Github => "github", - OAuthProviderEnum::Discord => "discord", - OAuthProviderEnum::Google => "google", - OAuthProviderEnum::Oidc => "oidc", - } - } -} /// Trait for all OAuth providers pub trait OAuthProvider: Send + Sync { @@ -158,7 +150,7 @@ impl<'a> OAuthService<'a> { db: &mut DbService, provider: OAuthProviderEnum, token: &StandardTokenResponse, - active_session: Option, + active_session: Option, ) -> AuthResult { // Get user info from provider let (oauth_client, oauth_provider) = self.oauth_provider(provider)?; diff --git a/server-new/src/services/auth/session.rs b/server-new/src/services/auth/session.rs index 35ce6cb..6bc93ee 100644 --- a/server-new/src/services/auth/session.rs +++ b/server-new/src/services/auth/session.rs @@ -9,7 +9,7 @@ use uuid::Uuid; use crate::{ db::{DbPool, DbService}, - extractors::session::{SessionMeta, UserSession}, + extractors::SessionMeta, services::auth::AuthResult, }; @@ -45,10 +45,10 @@ impl AuthSessionService { Ok(()) } - /// Extract the current user session if this is an active user session. - pub async fn user_session(&self, session: &Session) -> AuthResult> { + /// Extract the current user ID if this is an active user session. + pub async fn active_user_id(&self, session: &Session) -> AuthResult> { let user_id = session.get::(USER_ID_FIELD).await?; - Ok(user_id.map(UserSession::new)) + Ok(user_id) } /// Logout the user, deleting the current session. @@ -68,7 +68,7 @@ pub(super) fn user_id_from_record_data( data: &HashMap, ) -> StoreResult> { data.get(USER_ID_FIELD) - .map(|val| serde_json::from_value::(val.clone())) + .and_then(|val| val.as_str().map(Uuid::try_parse)) .transpose() .map_err(|_| StoreError::Encode("invalid user id field".to_owned())) } diff --git a/server-new/src/state.rs b/server-new/src/state.rs index c812e22..55471b9 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -32,7 +32,7 @@ pub struct AppStateInner { impl AppState { pub fn auth_service(&self) -> AuthService<'_> { - AuthService::new(&self.config, &self.oauth_providers) + AuthService::new(&self.config, &self.encryptor, &self.oauth_providers) } pub fn chat_service(&self) -> ChatService<'_> { ChatService::new(&self.db_pool, &self.tinistream) From cf001b8d4de09607e356d2a3b9f12258e718c13b Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sat, 4 Jul 2026 23:22:16 -0400 Subject: [PATCH 059/111] add proxy header auth --- server-new/config.toml | 6 ++ server-new/src/config.rs | 6 +- server-new/src/db/repositories/user.rs | 20 ++--- server-new/src/error.rs | 7 ++ server-new/src/extractors/user.rs | 44 +++++---- server-new/src/services/auth/api_key.rs | 9 +- server-new/src/services/auth/mod.rs | 8 +- server-new/src/services/auth/proxy.rs | 115 ++++++++++++++++++++++++ 8 files changed, 180 insertions(+), 35 deletions(-) create mode 100644 server-new/src/services/auth/proxy.rs diff --git a/server-new/config.toml b/server-new/config.toml index 3c97a52..1c6e7f5 100644 --- a/server-new/config.toml +++ b/server-new/config.toml @@ -21,6 +21,12 @@ streamer_api_key = "dev-streamer-api-key" cookie_name = "auth-rs-chat" session_length = 604800 +[auth.proxy] +enabled = false +username_header = "Remote-User" +name_header = "Remote-Name" +groups_header = "Remote-Groups" + [security] body_limit = 2097152 request_timeout = 120 diff --git a/server-new/src/config.rs b/server-new/src/config.rs index aa54d40..5b3181c 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -6,7 +6,10 @@ use figment::providers::{Env, Format, Toml}; use serde::Deserialize; use crate::{ - services::auth::oauth::{DiscordOAuthConfig, GitHubOAuthConfig, GoogleOAuthConfig, OidcConfig}, + services::auth::{ + oauth::{DiscordOAuthConfig, GitHubOAuthConfig, GoogleOAuthConfig, OidcConfig}, + proxy::ProxyHeaderConfig, + }, state::AppState, }; @@ -51,6 +54,7 @@ pub struct AuthConfig { pub discord: Option, pub google: Option, pub oidc: Option, + pub proxy: ProxyHeaderConfig, } #[derive(Debug, Clone, Deserialize)] diff --git a/server-new/src/db/repositories/user.rs b/server-new/src/db/repositories/user.rs index abb4908..25f3908 100644 --- a/server-new/src/db/repositories/user.rs +++ b/server-new/src/db/repositories/user.rs @@ -81,16 +81,16 @@ impl<'a> UserRepository<'a> { Ok(user) } - // pub async fn find_by_sso_username(&mut self, username: &str) -> Result, Error> { - // let user_id = users::table - // .filter(users::sso_username.eq(username)) - // .select(users::id) - // .first(self.db) - // .await - // .optional()?; - - // Ok(user_id) - // } + pub async fn find_by_sso_username(&mut self, username: &str) -> Result, Error> { + let user_id = users::table + .filter(users::sso_username.eq(username)) + .select(users::id) + .first(self.db) + .await + .optional()?; + + Ok(user_id) + } pub async fn update( &mut self, diff --git a/server-new/src/error.rs b/server-new/src/error.rs index 3198802..a9dce07 100644 --- a/server-new/src/error.rs +++ b/server-new/src/error.rs @@ -5,6 +5,8 @@ use axum::{ }; use serde::Serialize; +use crate::db::DbPoolError; + /// Global result type that can be used for API route handlers pub type AppResult = Result; @@ -55,6 +57,11 @@ impl From for AppError { Self::internal(anyhow::Error::from(err).context("database error")) } } +impl From for AppError { + fn from(err: DbPoolError) -> Self { + Self::internal(anyhow::Error::from(err).context("database pool error")) + } +} #[derive(Debug, Serialize)] struct ErrorResponse { diff --git a/server-new/src/extractors/user.rs b/server-new/src/extractors/user.rs index c7200e9..427cff0 100644 --- a/server-new/src/extractors/user.rs +++ b/server-new/src/extractors/user.rs @@ -9,7 +9,7 @@ use uuid::Uuid; use crate::{db::DbService, error::AppError, state::AppState}; /** -Represents an active user, extracted from the session or API key. This can be used +Represents an active user, extracted from the session, proxy headers, or API key. This can be used as an extractor in route handlers: - If used as `CurrentUser`, request will automatically return an unauthorized error if there is no active user. @@ -34,34 +34,44 @@ impl OptionalFromRequestParts for CurrentUser { parts: &mut axum::http::request::Parts, state: &AppState, ) -> Result, Self::Rejection> { + let auth_service = state.auth_service(); + // If there is an Authorization header, validate the API key if let Some(auth_header) = parts .headers .get(header::AUTHORIZATION) .and_then(|bytes| bytes.to_str().ok()) { - let mut db = DbService::from_pool(&state.db_pool) - .await - .map_err(|err| AppError::internal(err.into()))?; - let user_id = state - .auth_service() + let user_id = auth_service .api_keys() - .validate_api_key(&mut db, auth_header) + .validate_api_key(&state.db_pool, auth_header) .await?; Ok(Some(Self { user_id })) - } - // Check for session - else { + } else { + // If SSO header / proxy auth is enabled, check forwarded headers first + if state.config.auth.proxy.enabled { + let proxy_service = auth_service.proxy(); + if let Some(proxy_user) = proxy_service.extract_proxy_user(&parts.headers)? { + let mut db = DbService::from_pool(&state.db_pool).await?; + match proxy_service.find_proxy_user(&mut db, &proxy_user).await? { + Some(user_id) => return Ok(Some(Self::new(user_id))), + None => { + let new_user = proxy_service + .create_proxy_user(&mut db, &proxy_user) + .await?; + return Ok(Some(Self::new(new_user.id))); + } + } + } + } + + // Check for session let session = parts .extensions .get::() .ok_or_else(|| AppError::internal(anyhow!("session not attached to request")))?; - let maybe_user_id = state - .auth_service() - .session() - .active_user_id(&session) - .await?; + let maybe_user_id = auth_service.session().active_user_id(&session).await?; Ok(maybe_user_id.map(Self::new)) } @@ -77,8 +87,8 @@ impl FromRequestParts for CurrentUser { ) -> Result { match >::from_request_parts(parts, state).await? { - Some(user_session) => Ok(user_session), - None => Err(AppError::unauthorized("no active session")), + Some(current_user) => Ok(current_user), + None => Err(AppError::unauthorized("no active user")), } } } diff --git a/server-new/src/services/auth/api_key.rs b/server-new/src/services/auth/api_key.rs index ae79d90..88cda9f 100644 --- a/server-new/src/services/auth/api_key.rs +++ b/server-new/src/services/auth/api_key.rs @@ -1,7 +1,7 @@ use uuid::Uuid; use crate::{ - db::{DbService, models::NewChatRsApiKey}, + db::{DbPool, DbService, models::NewChatRsApiKey}, services::auth::{ encryption::Encryptor, error::{AuthError, AuthResult}, @@ -49,11 +49,7 @@ impl<'r> ApiKeyService<'r> { } /// Validate the API key and get the user ID - pub async fn validate_api_key( - &self, - db: &mut DbService, - auth_header: &str, - ) -> AuthResult { + pub async fn validate_api_key(&self, db: &DbPool, auth_header: &str) -> AuthResult { let (nonce, ciphertext) = auth_header .strip_prefix(API_KEY_HEADER_PREFIX) .and_then(|s| s.split_once('|')) @@ -72,6 +68,7 @@ impl<'r> ApiKeyService<'r> { .map_err(|_| AuthError::Unauthorized("couldn't parse API key id")) })?; + let mut db = DbService::from_pool(db).await?; match db.api_keys().find_by_id(&api_key_id).await? { Some(api_key) => Ok(api_key.user_id), None => Err(AuthError::Unauthorized("API key not found")), diff --git a/server-new/src/services/auth/mod.rs b/server-new/src/services/auth/mod.rs index bda9f1a..c296c18 100644 --- a/server-new/src/services/auth/mod.rs +++ b/server-new/src/services/auth/mod.rs @@ -9,6 +9,7 @@ pub mod api_key; pub mod encryption; mod error; pub mod oauth; +pub mod proxy; pub mod session; pub mod session_store; @@ -52,8 +53,13 @@ impl<'r> AuthService<'r> { oauth::OAuthService::new(self.config, self.oauth_providers) } + /// Access proxy auth functions + pub fn proxy(&self) -> proxy::ProxyService<'r> { + proxy::ProxyService::new(&self.config.auth.proxy) + } + /// Access API key functions - pub fn api_keys(self) -> api_key::ApiKeyService<'r> { + pub fn api_keys(&self) -> api_key::ApiKeyService<'r> { api_key::ApiKeyService::new(self.encryptor) } } diff --git a/server-new/src/services/auth/proxy.rs b/server-new/src/services/auth/proxy.rs new file mode 100644 index 0000000..45f5383 --- /dev/null +++ b/server-new/src/services/auth/proxy.rs @@ -0,0 +1,115 @@ +use axum::http::HeaderMap; +use serde::Deserialize; +use uuid::Uuid; + +use crate::{ + db::{ + DbService, + models::{ChatRsUser, NewChatRsUser}, + }, + services::auth::error::{AuthError, AuthResult}, +}; + +/// SSO / forward auth proxy header configuration +#[derive(Debug, Clone, Deserialize)] +pub struct ProxyHeaderConfig { + /// Whether proxy header authentication is enabled + pub enabled: bool, + /// Header for unique, identifying username (default: `Remote-User`) + username_header: String, + /// Header for display name (default: `Remote-Name`) + name_header: String, + /// Header for space-delimited groups/roles of the user (default: `Remote-Groups`) + groups_header: String, + /// If set, only users in these groups will be allowed to access the app + user_groups: Option>, + /// URL to redirect to in order to log out of the remote service + logout_url: Option, +} + +impl Default for ProxyHeaderConfig { + fn default() -> Self { + Self { + enabled: false, + username_header: String::from("Remote-User"), + name_header: String::from("Remote-Name"), + groups_header: String::from("Remote-Groups"), + user_groups: None, + logout_url: None, + } + } +} + +pub struct ProxyService<'r> { + config: &'r ProxyHeaderConfig, +} + +pub struct ProxyUser { + username: String, + name: String, +} + +impl<'r> ProxyService<'r> { + pub fn new(config: &'r ProxyHeaderConfig) -> Self { + Self { config } + } + + pub fn extract_proxy_user(&self, headers: &HeaderMap) -> AuthResult> { + let Some(username) = headers.get(&self.config.username_header) else { + return Ok(None); + }; + let name = headers.get(&self.config.name_header).unwrap_or(username); + let groups = headers + .get(&self.config.groups_header) + .and_then(|groups| groups.to_str().ok()) + .unwrap_or_default(); + + if let Some(ref allowed_groups) = self.config.user_groups { + if !is_proxy_user_allowed(groups, allowed_groups) { + return Err(AuthError::Unauthorized("proxy user not in allowed group")); + } + } + + Ok(Some(ProxyUser { + username: username.to_str().unwrap_or_default().to_owned(), + name: name.to_str().unwrap_or_default().to_owned(), + })) + } + + pub async fn find_proxy_user( + &self, + db: &mut DbService, + proxy_user: &ProxyUser, + ) -> AuthResult> { + let user_id = db + .users() + .find_by_sso_username(&proxy_user.username) + .await?; + Ok(user_id) + } + + pub async fn create_proxy_user( + &self, + db: &mut DbService, + proxy_user: &ProxyUser, + ) -> AuthResult { + let new_user = db + .users() + .create(NewChatRsUser { + sso_username: Some(&proxy_user.username), + name: &proxy_user.name, + ..Default::default() + }) + .await?; + Ok(new_user) + } +} + +fn is_proxy_user_allowed(user_groups: &str, allowed_groups: &[String]) -> bool { + for user_group in user_groups.split(' ') { + if allowed_groups.iter().any(|g| g == user_group) { + return true; + } + } + return false; +} From 224c42a462ff319f9ec4e29eb39fa2c9af81ade4 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sun, 5 Jul 2026 01:02:33 -0400 Subject: [PATCH 060/111] extract config defaults --- examples/compose/compose.yml | 1 + examples/compose/config.toml | 40 +++++++ server-new/.env.example | 3 +- server-new/Dockerfile | 14 +-- server-new/config.toml | 21 +--- server-new/src/config.rs | 107 ++++++++++++++---- server-new/src/services/auth/oauth/discord.rs | 4 +- server-new/src/services/auth/oauth/github.rs | 4 +- server-new/src/services/auth/oauth/google.rs | 4 +- server-new/src/services/auth/oauth/oidc.rs | 4 +- server-new/src/services/auth/proxy.rs | 4 +- 11 files changed, 141 insertions(+), 65 deletions(-) create mode 100644 examples/compose/compose.yml create mode 100644 examples/compose/config.toml diff --git a/examples/compose/compose.yml b/examples/compose/compose.yml new file mode 100644 index 0000000..4640904 --- /dev/null +++ b/examples/compose/compose.yml @@ -0,0 +1 @@ +# TODO diff --git a/examples/compose/config.toml b/examples/compose/config.toml new file mode 100644 index 0000000..2509736 --- /dev/null +++ b/examples/compose/config.toml @@ -0,0 +1,40 @@ +[server] +host = "0.0.0.0" +port = 8080 +base_url = "https://example.com" # set to domain & path where you're hosting +log_level = "info" +request_id_header = "x-request-id" + +[database] +url = "postgres://user:pass@mydb:5432" + +[redis] +url = "redis://pass@myredis:6379" +pool_size = 4 +timeout = 10 + +[services] +streamer_url = "http://tinistream:8081" +streamer_api_key = "tinistream-api-key" + +[auth] +cookie_name = "auth-rs-chat" +session_length = 604800 # seconds + +# Configure GitHub / Discord / Google login +[auth.github] +client_id = "" +client_secret = "" + +# Configure forward header authentication +# ⚠️ Only use with a secure proxy like Authelia or Tinyauth +[auth.proxy] +enabled = false +username_header = "Remote-User" +name_header = "Remote-Name" +groups_header = "Remote-Groups" + +# Configure various security settings +[security] +body_limit = 2097152 # bytes +request_timeout = 120 # seconds diff --git a/server-new/.env.example b/server-new/.env.example index a2962b6..3120ed4 100644 --- a/server-new/.env.example +++ b/server-new/.env.example @@ -1,5 +1,4 @@ -# Database (2nd one is needed for Diesel CLI) -RS_CHAT_DATABASE__URL=postgres://postgres:postgres@localhost/postgres +# Database (needed for Diesel CLI) DATABASE_URL=postgres://postgres:postgres@localhost/postgres # Auth diff --git a/server-new/Dockerfile b/server-new/Dockerfile index 676bffc..bc26e19 100644 --- a/server-new/Dockerfile +++ b/server-new/Dockerfile @@ -25,22 +25,14 @@ RUN --mount=type=cache,id=rust_target,target=/app/target \ ### Run server ### FROM debian:${DEBIAN_VERSION}-slim AS run -# Create non-root user -ARG UID=10001 -RUN adduser \ - --disabled-password \ - --gecos "" \ - --home "/home/appuser" \ - --shell "/sbin/nologin" \ - --uid "${UID}" \ - appuser -USER appuser +RUN apt-get update -qq && \ + apt-get install ca-certificates -qq -y && \ + apt-get clean # Copy server binary COPY --from=build --chown=appuser /app/run-server /usr/local/bin/ # Run server WORKDIR /app -COPY --chown=appuser config.toml ./config.toml ENV RS_CHAT_SERVER__HOST=0.0.0.0 CMD ["run-server"] diff --git a/server-new/config.toml b/server-new/config.toml index 1c6e7f5..c5bc5ae 100644 --- a/server-new/config.toml +++ b/server-new/config.toml @@ -1,32 +1,17 @@ +# Configuration for local development + [server] host = "127.0.0.1" port = 8080 base_url = "http://localhost:8080" log_level = "info" -request_id_header = "x-request-id" [database] -url = "postgres://localhost" +url = "postgres://postgres:postgres@localhost/postgres" [redis] url = "redis://localhost:6379" -pool_size = 4 -timeout = 10 [services] streamer_url = "http://localhost:8081" streamer_api_key = "dev-streamer-api-key" - -[auth] -cookie_name = "auth-rs-chat" -session_length = 604800 - -[auth.proxy] -enabled = false -username_header = "Remote-User" -name_header = "Remote-Name" -groups_header = "Remote-Groups" - -[security] -body_limit = 2097152 -request_timeout = 120 diff --git a/server-new/src/config.rs b/server-new/src/config.rs index 5b3181c..29788ba 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -2,8 +2,8 @@ use std::net::IpAddr; use anyhow::Context; use axum_plugin::AdHocPlugin; -use figment::providers::{Env, Format, Toml}; -use serde::Deserialize; +use figment::providers::{Env, Format, Serialized, Toml}; +use serde::{Deserialize, Serialize}; use crate::{ services::auth::{ @@ -13,8 +13,22 @@ use crate::{ state::AppState, }; +/// Plugin that reads and validates configuration, and adds it to server state +pub fn plugin() -> AdHocPlugin { + AdHocPlugin::named("Config").on_init(async |mut state| { + let config = extract_config()?; + tracing::info!( + log_level = config.server.log_level, + base_url = config.server.base_url, + "Config loaded!" + ); + state.insert(config); + Ok(state) + }) +} + /// Parsed app configuration -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Default, Clone, Serialize, Deserialize)] pub struct AppConfig { pub server: ServerConfig, pub database: DatabaseConfig, @@ -24,7 +38,7 @@ pub struct AppConfig { pub redis: RedisConfig, } -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct ServerConfig { pub host: IpAddr, pub port: u16, @@ -33,19 +47,46 @@ pub struct ServerConfig { pub request_id_header: String, pub ip_header: Option, } +impl Default for ServerConfig { + fn default() -> Self { + Self { + host: IpAddr::V4(std::net::Ipv4Addr::LOCALHOST), + port: 8080, + base_url: String::from("http://localhost:8080"), + log_level: String::from("info"), + request_id_header: String::from("x-request-id"), + ip_header: None, + } + } +} -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct DatabaseConfig { pub url: String, } +impl Default for DatabaseConfig { + fn default() -> Self { + Self { + url: String::from("postgres://localhost:5432"), + } + } +} -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct ServiceConfig { pub streamer_url: String, pub streamer_api_key: String, } +impl Default for ServiceConfig { + fn default() -> Self { + Self { + streamer_url: String::from("http://localhost:8081"), + streamer_api_key: String::new(), + } + } +} -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct AuthConfig { pub encryption_key: String, pub cookie_name: String, @@ -56,37 +97,55 @@ pub struct AuthConfig { pub oidc: Option, pub proxy: ProxyHeaderConfig, } +impl Default for AuthConfig { + fn default() -> Self { + Self { + encryption_key: String::new(), + cookie_name: String::from("auth-rs-chat"), + session_length: 604800, // 1 week in seconds + github: None, + discord: None, + google: None, + oidc: None, + proxy: ProxyHeaderConfig::default(), + } + } +} -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct SecurityConfig { pub body_limit: usize, pub request_timeout: u64, } +impl Default for SecurityConfig { + fn default() -> Self { + Self { + body_limit: 2097152, // 2 MB + request_timeout: 120, // 2 minutes + } + } +} -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct RedisConfig { pub url: String, pub pool_size: usize, pub timeout: u64, } - -/// Plugin that reads and validates configuration, and adds it to server state -pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("Config").on_init(async |mut state| { - let config = extract_config()?; - tracing::info!( - log_level = config.server.log_level, - base_url = config.server.base_url, - "Config loaded!" - ); - state.insert(config); - Ok(state) - }) +impl Default for RedisConfig { + fn default() -> Self { + Self { + url: String::from("redis://localhost:6379"), + pool_size: 4, + timeout: 10, // 10 seconds + } + } } -/// Extract configuration from config.toml, then environment variables. +/// Extract configuration from defaults, local `config.toml`, then `RS_CHAT_` environment variables. +/// See https://docs.rs/figment/latest/figment/index.html#for-application-authors fn extract_config() -> anyhow::Result { - let config = figment::Figment::new() + let config = figment::Figment::from(Serialized::defaults(AppConfig::default())) .merge(Toml::file("config.toml")) .merge(Env::prefixed("RS_CHAT_").split("__")) .extract::() diff --git a/server-new/src/services/auth/oauth/discord.rs b/server-new/src/services/auth/oauth/discord.rs index c803803..d6196df 100644 --- a/server-new/src/services/auth/oauth/discord.rs +++ b/server-new/src/services/auth/oauth/discord.rs @@ -1,5 +1,5 @@ use futures::future::BoxFuture; -use serde::Deserialize; +use serde::{Deserialize, Serialize}; use simple_oauth::types::{OAuthCredentials, UserInfo}; use crate::{ @@ -7,7 +7,7 @@ use crate::{ services::auth::{AuthResult, oauth::OAuthProvider}, }; -#[derive(Clone, Debug, Deserialize)] +#[derive(Clone, Debug, Serialize, Deserialize)] pub struct DiscordOAuthConfig { client_id: u64, client_secret: String, diff --git a/server-new/src/services/auth/oauth/github.rs b/server-new/src/services/auth/oauth/github.rs index 092bb7a..35314a7 100644 --- a/server-new/src/services/auth/oauth/github.rs +++ b/server-new/src/services/auth/oauth/github.rs @@ -1,5 +1,5 @@ use futures::future::BoxFuture; -use serde::Deserialize; +use serde::{Deserialize, Serialize}; use simple_oauth::{ SimpleOAuthProvider, types::{OAuthCredentials, UserInfo}, @@ -10,7 +10,7 @@ use crate::{ services::auth::{AuthResult, oauth::OAuthProvider}, }; -#[derive(Clone, Debug, Deserialize)] +#[derive(Clone, Debug, Serialize, Deserialize)] pub struct GitHubOAuthConfig { client_id: String, client_secret: String, diff --git a/server-new/src/services/auth/oauth/google.rs b/server-new/src/services/auth/oauth/google.rs index 16d7520..a81c8bb 100644 --- a/server-new/src/services/auth/oauth/google.rs +++ b/server-new/src/services/auth/oauth/google.rs @@ -1,5 +1,5 @@ use futures::future::BoxFuture; -use serde::Deserialize; +use serde::{Deserialize, Serialize}; use simple_oauth::types::{OAuthCredentials, UserInfo}; use crate::{ @@ -7,7 +7,7 @@ use crate::{ services::auth::{AuthResult, oauth::OAuthProvider}, }; -#[derive(Clone, Debug, Deserialize)] +#[derive(Clone, Debug, Serialize, Deserialize)] pub struct GoogleOAuthConfig { client_id: String, client_secret: String, diff --git a/server-new/src/services/auth/oauth/oidc.rs b/server-new/src/services/auth/oauth/oidc.rs index 6d4eef6..efbfec0 100644 --- a/server-new/src/services/auth/oauth/oidc.rs +++ b/server-new/src/services/auth/oauth/oidc.rs @@ -1,5 +1,5 @@ use futures::future::BoxFuture; -use serde::Deserialize; +use serde::{Deserialize, Serialize}; use simple_oauth::{ SimpleOAuthProvider, types::{OAuthCredentials, OidcDiscovery}, @@ -12,7 +12,7 @@ use crate::{ use super::OAuthProvider; -#[derive(Clone, Debug, Deserialize)] +#[derive(Clone, Debug, Serialize, Deserialize)] pub struct OidcConfig { pub name: Option, client_id: String, diff --git a/server-new/src/services/auth/proxy.rs b/server-new/src/services/auth/proxy.rs index 45f5383..c358ba3 100644 --- a/server-new/src/services/auth/proxy.rs +++ b/server-new/src/services/auth/proxy.rs @@ -1,5 +1,5 @@ use axum::http::HeaderMap; -use serde::Deserialize; +use serde::{Deserialize, Serialize}; use uuid::Uuid; use crate::{ @@ -11,7 +11,7 @@ use crate::{ }; /// SSO / forward auth proxy header configuration -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct ProxyHeaderConfig { /// Whether proxy header authentication is enabled pub enabled: bool, From bea7d131b50c0b5b534d6b739e78396f1227e78c Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sun, 5 Jul 2026 21:34:33 -0400 Subject: [PATCH 061/111] api doc - checkpoint --- server-new/Cargo.lock | 326 ++++++++++++++---- server-new/Cargo.toml | 11 +- server-new/Dockerfile | 2 +- server-new/src/api/api_key.rs | 99 ++++-- server-new/src/api/auth.rs | 125 ++++--- server-new/src/api/chat.rs | 93 +++-- server-new/src/api/health.rs | 11 +- server-new/src/api/mod.rs | 64 ++-- server-new/src/db/models/api_key.rs | 4 +- server-new/src/db/models/user.rs | 4 +- server-new/src/error.rs | 7 +- server-new/src/extractors/auth_config.rs | 7 +- server-new/src/extractors/database.rs | 2 + server-new/src/extractors/mod.rs | 2 +- server-new/src/extractors/session.rs | 34 +- server-new/src/extractors/user.rs | 3 +- server-new/src/llm/types.rs | 3 +- server-new/src/services/auth/oauth.rs | 14 +- server-new/src/services/auth/session.rs | 19 +- server-new/src/services/auth/session_store.rs | 7 +- server-new/src/services/chat/mod.rs | 20 +- 21 files changed, 605 insertions(+), 252 deletions(-) diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index da4ef72..f6ebf95 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -81,6 +81,59 @@ dependencies = [ "memchr", ] +[[package]] +name = "aide" +version = "0.15.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6966317188cdfe54c58c0900a195d021294afb3ece9b7073d09e4018dbb1e3a2" +dependencies = [ + "axum", + "bytes", + "cfg-if", + "http", + "indexmap", + "schemars 0.9.0", + "serde", + "serde_json", + "serde_qs 0.14.0", + "thiserror 2.0.18", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "aide" +version = "0.16.0-alpha.4" +source = "git+https://github.com/hniksic/aide.git?rev=7246c20#7246c20903ce9e87b768917589a7568663a1eda3" +dependencies = [ + "aide-macros", + "axum", + "bytes", + "cfg-if", + "http", + "indexmap", + "schemars 1.2.1", + "serde", + "serde_json", + "serde_qs 1.1.2", + "thiserror 2.0.18", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "aide-macros" +version = "0.16.0-alpha.4" +source = "git+https://github.com/hniksic/aide.git?rev=7246c20#7246c20903ce9e87b768917589a7568663a1eda3" +dependencies = [ + "darling 0.23.0", + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "android_system_properties" version = "0.1.5" @@ -249,6 +302,31 @@ dependencies = [ "tracing", ] +[[package]] +name = "axum-extra" +version = "0.12.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "be44683b41ccb9ab2d23a5230015c9c3c55be97a25e4428366de8873103f7970" +dependencies = [ + "axum", + "axum-core", + "bytes", + "form_urlencoded", + "futures-core", + "futures-util", + "http", + "http-body", + "http-body-util", + "mime", + "pin-project-lite", + "serde_core", + "serde_html_form", + "serde_path_to_error", + "tower-layer", + "tower-service", + "tracing", +] + [[package]] name = "axum-helmet" version = "1.0.2" @@ -262,6 +340,17 @@ dependencies = [ "tower-service", ] +[[package]] +name = "axum-macros" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7aa268c23bfbbd2c4363b9cd302a4f504fb2a9dfe7e3451d66f35dd392e20aca" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "axum-plugin" version = "0.2.0" @@ -284,6 +373,30 @@ dependencies = [ "syn", ] +[[package]] +name = "axum-typed-routing" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a61ccff37342bbcd7fc468ab6d96e9547a3c81bb2b027f108dda07e0757ab477" +dependencies = [ + "aide 0.15.1", + "axum", + "axum-extra", + "axum-macros", + "axum-typed-routing-macros", +] + +[[package]] +name = "axum-typed-routing-macros" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "299699055db19bfe910cb2275056062b6dbb0198f7d6d365e1f4f5e2f9f6138f" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "base64" version = "0.22.1" @@ -692,6 +805,27 @@ dependencies = [ "serde_core", ] +[[package]] +name = "derive_more" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d751e9e49156b02b44f9c1815bcb94b984cdcc4396ecc32521c739452808b134" +dependencies = [ + "derive_more-impl", +] + +[[package]] +name = "derive_more-impl" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "799a97264921d8623a957f6c3b9011f3b5492f557bbb7a5a19b7fa6d06ba8dcb" +dependencies = [ + "proc-macro2", + "quote", + "rustc_version", + "syn", +] + [[package]] name = "diesel" version = "2.3.10" @@ -832,6 +966,12 @@ version = "1.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "92773504d58c093f6de2459af4af33faa518c13451eb8f2b5698ed3d36e7c813" +[[package]] +name = "dyn-clone" +version = "1.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d0881ea181b1df73ff77ffaaf9c7544ecc11e82fba9b5f27b262a3c73a332555" + [[package]] name = "either" version = "1.16.0" @@ -1776,12 +1916,6 @@ dependencies = [ "windows-link", ] -[[package]] -name = "paste" -version = "1.0.15" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" - [[package]] name = "pear" version = "0.2.9" @@ -2124,6 +2258,26 @@ dependencies = [ "bitflags", ] +[[package]] +name = "ref-cast" +version = "1.0.25" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f354300ae66f76f1c85c5f84693f0ce81d747e2c3f21a45fef496d89c960bf7d" +dependencies = [ + "ref-cast-impl", +] + +[[package]] +name = "ref-cast-impl" +version = "1.0.25" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7186006dcb21920990093f30e3dea63b7d6e977bf1256be20c3563a5db070da" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "regex-automata" version = "0.4.14" @@ -2240,13 +2394,16 @@ name = "rs-chat-api" version = "0.1.0" dependencies = [ "aes-gcm 0.11.0", + "aide 0.16.0-alpha.4", "anyhow", "async-stream", "async-trait", "axum", "axum-helmet", "axum-plugin", + "axum-typed-routing", "chrono", + "derive_more", "diesel", "diesel-async", "diesel-derive-enum", @@ -2259,6 +2416,7 @@ dependencies = [ "hex", "reqwest", "reqwest-websocket", + "schemars 1.2.1", "serde", "serde_json", "serde_with", @@ -2276,9 +2434,6 @@ dependencies = [ "tracing", "tracing-appender", "tracing-subscriber", - "utoipa", - "utoipa-axum", - "utoipa-scalar", "uuid", ] @@ -2402,6 +2557,60 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "schemars" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4cd191f9397d57d581cddd31014772520aa448f65ef991055d7f61582c65165f" +dependencies = [ + "dyn-clone", + "indexmap", + "ref-cast", + "schemars_derive 0.9.0", + "serde", + "serde_json", +] + +[[package]] +name = "schemars" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a2b42f36aa1cd011945615b92222f6bf73c599a102a300334cd7f8dbeec726cc" +dependencies = [ + "chrono", + "dyn-clone", + "indexmap", + "ref-cast", + "schemars_derive 1.2.1", + "serde", + "serde_json", + "uuid", +] + +[[package]] +name = "schemars_derive" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5016d94c77c6d32f0b8e08b781f7dc8a90c2007d4e77472cc2807bc10a8438fe" +dependencies = [ + "proc-macro2", + "quote", + "serde_derive_internals", + "syn", +] + +[[package]] +name = "schemars_derive" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7d115b50f4aaeea07e79c1912f645c7513d81715d0420f8bc77a18c6260b307f" +dependencies = [ + "proc-macro2", + "quote", + "serde_derive_internals", + "syn", +] + [[package]] name = "scopeguard" version = "1.2.0" @@ -2467,6 +2676,30 @@ dependencies = [ "syn", ] +[[package]] +name = "serde_derive_internals" +version = "0.29.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "18d26a20a969b9e3fdf2fc2d9f21eda6c40e2de84c9408bb5d3b05d499aae711" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "serde_html_form" +version = "0.2.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b2f2d7ff8a2140333718bb329f5c40fc5f0865b84c426183ce14c97d2ab8154f" +dependencies = [ + "form_urlencoded", + "indexmap", + "itoa", + "ryu", + "serde_core", +] + [[package]] name = "serde_json" version = "1.0.150" @@ -2491,6 +2724,32 @@ dependencies = [ "serde_core", ] +[[package]] +name = "serde_qs" +version = "0.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b417bedc008acbdf6d6b4bc482d29859924114bbe2650b7921fb68a261d0aa6" +dependencies = [ + "axum", + "futures", + "percent-encoding", + "serde", + "thiserror 2.0.18", +] + +[[package]] +name = "serde_qs" +version = "1.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "67d525c8ff68aa99e5818302259bdd02d86d0303710616f39c0f44846ff6d332" +dependencies = [ + "axum", + "itoa", + "percent-encoding", + "ryu", + "serde", +] + [[package]] name = "serde_spanned" version = "0.6.9" @@ -3389,55 +3648,6 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" -[[package]] -name = "utoipa" -version = "5.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8bde15df68e80b16c7d16b9616e80770ad158988daa56a27dccd1e55558b0160" -dependencies = [ - "indexmap", - "serde", - "serde_json", - "utoipa-gen", -] - -[[package]] -name = "utoipa-axum" -version = "0.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7c25bae5bccc842449ec0c5ddc5cbb6a3a1eaeac4503895dc105a1138f8234a0" -dependencies = [ - "axum", - "paste", - "tower-layer", - "tower-service", - "utoipa", -] - -[[package]] -name = "utoipa-gen" -version = "5.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6ba0b99ee52df3028635d93840c797102da61f8a7bb3cf751032455895b52ef8" -dependencies = [ - "proc-macro2", - "quote", - "syn", - "uuid", -] - -[[package]] -name = "utoipa-scalar" -version = "0.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "59559e1509172f6b26c1cdbc7247c4ddd1ac6560fe94b584f81ee489b141f719" -dependencies = [ - "axum", - "serde", - "serde_json", - "utoipa", -] - [[package]] name = "uuid" version = "1.23.3" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 6355bf5..3db8c8a 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -7,6 +7,11 @@ publish = false [dependencies] aes-gcm = "0.11.0" +aide = { + git = "https://github.com/hniksic/aide.git", + rev = "7246c20", + features = ["axum", "axum-json", "axum-query", "macros", "swagger"] +} anyhow = "1.0.102" async-stream = "0.3.6" async-trait = "0.1.89" @@ -16,11 +21,13 @@ axum-plugin = { git = "https://git.fasharp.io/fa-sharp/axum-plugin", rev = "9f72278b3c" } +axum-typed-routing = { version = "0.4.0", features = ["aide"] } chrono = { version = "0.4.45", default-features = false, features = ["now", "serde", "std"] } +derive_more = { version = "2.1.1", features = ["deref", "into"] } diesel = { version = "2.3.10", default-features = false, @@ -48,6 +55,7 @@ reqwest = { features = ["default-tls", "json", "stream"] } reqwest-websocket = { version = "0.6.0", features = ["json"] } +schemars = { version = "1.2.1", features = ["chrono04", "uuid1"] } serde = { version = "1.0.228", features = ["derive"] } serde_json = "1.0.150" serde_with = { @@ -90,7 +98,4 @@ tower-sessions-redis-store = { tracing = "0.1.44" tracing-appender = "0.2.5" tracing-subscriber = { version = "0.3.23", features = ["env-filter", "json"] } -utoipa = { version = "5.5.0", features = ["chrono", "uuid"] } -utoipa-axum = "0.2.0" -utoipa-scalar = { version = "0.3.0", features = ["axum"] } uuid = { version = "1.23.3", features = ["serde", "v4"] } diff --git a/server-new/Dockerfile b/server-new/Dockerfile index bc26e19..fe58b19 100644 --- a/server-new/Dockerfile +++ b/server-new/Dockerfile @@ -30,7 +30,7 @@ RUN apt-get update -qq && \ apt-get clean # Copy server binary -COPY --from=build --chown=appuser /app/run-server /usr/local/bin/ +COPY --from=build /app/run-server /usr/local/bin/ # Run server WORKDIR /app diff --git a/server-new/src/api/api_key.rs b/server-new/src/api/api_key.rs index 6417d97..61ed15e 100644 --- a/server-new/src/api/api_key.rs +++ b/server-new/src/api/api_key.rs @@ -1,11 +1,9 @@ -use axum::{ - Json, - extract::{Path, State}, - http::StatusCode, -}; +use aide::{OperationInput, OperationIo, axum::ApiRouter}; +use axum::{Json, extract::State, http::StatusCode}; +use axum_typed_routing::{TypedApiRouter, api_route}; +use derive_more::Deref; +use schemars::JsonSchema; use serde::{Deserialize, Serialize}; -use utoipa::ToSchema; -use utoipa_axum::{router::OpenApiRouter, routes}; use uuid::Uuid; use crate::{ @@ -16,16 +14,17 @@ use crate::{ state::AppState, }; -pub fn routes() -> OpenApiRouter { - OpenApiRouter::new().routes(routes!(list_api_keys, create_api_key, delete_api_key)) +pub fn routes() -> ApiRouter { + ApiRouter::new() + .typed_api_route(list_api_keys) + .typed_api_route(create_api_key) + .typed_api_route(delete_api_key) } -/// List all API keys -#[utoipa::path( - get, path = "", - responses((status = OK, body = Vec)), - tag = ApiTag::ApiKey.into()) -] +#[api_route(GET "/" with AppState { + summary: "List API keys", + transform: |op| op.tag(ApiTag::ApiKey.into()), +})] async fn list_api_keys( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, @@ -34,24 +33,21 @@ async fn list_api_keys( Ok(Json(keys)) } -#[derive(Deserialize, ToSchema)] +#[derive(Deserialize, JsonSchema)] struct ApiKeyCreateInput { name: String, } -#[derive(Serialize, ToSchema)] +#[derive(Serialize, JsonSchema)] struct ApiKeyCreateResponse { id: Uuid, key: String, } -/// Create an API key -#[utoipa::path( - post, path = "", - request_body = ApiKeyCreateInput, - responses((status = OK, body = ApiKeyCreateResponse)), - tag = ApiTag::ApiKey.into()) -] +#[api_route(POST "/" { + summary: "Create API key", + transform: |op| op.tag(ApiTag::ApiKey.into()), +})] async fn create_api_key( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, @@ -67,19 +63,56 @@ async fn create_api_key( Ok(Json(ApiKeyCreateResponse { id, key })) } -/// Delete an API key -#[utoipa::path( - delete, path = "/{id}", - params(("id" = Uuid, Path)), - responses((status = NO_CONTENT)), - tag = ApiTag::ApiKey.into()) -] +#[derive(Deref, Deserialize, OperationIo, JsonSchema)] +pub struct ApiKeyPath { + pub id: Uuid, +} + +#[derive(Deref, Serialize, Deserialize, JsonSchema)] +pub struct UuidPath(pub Uuid); +impl OperationInput for UuidPath { + fn operation_input( + ctx: &mut aide::generate::GenContext, + operation: &mut aide::openapi::Operation, + ) { + use aide::openapi::{ + Parameter, ParameterData, ParameterSchemaOrContent, ReferenceOr, SchemaObject, + }; + + operation + .parameters + .push(ReferenceOr::Item(Parameter::Path { + parameter_data: ParameterData { + name: UuidPath::schema_name().into(), + description: None, + required: true, + deprecated: Default::default(), + format: ParameterSchemaOrContent::Schema(SchemaObject { + json_schema: UuidPath::json_schema(&mut ctx.schema), + external_docs: None, + example: None, + }), + example: Default::default(), + examples: Default::default(), + explode: Default::default(), + extensions: Default::default(), + }, + style: aide::openapi::PathStyle::Simple, + })) + } +} + +#[api_route(DELETE "/{id}" with AppState { + summary: "Delete API key", + responses: { 204: () }, + transform: |op| op.tag(ApiTag::ApiKey.into()), +})] async fn delete_api_key( + id: UuidPath, CurrentUser { user_id }: CurrentUser, - Path(api_key_id): Path, Database(mut db): Database, ) -> AppResult { - let _deleted_id = db.api_keys().delete(&user_id, &api_key_id).await?; + let _ = db.api_keys().delete(&user_id, &id).await?; Ok(StatusCode::NO_CONTENT) } diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index 9652f24..63e2ec9 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -1,138 +1,131 @@ +use aide::{ + axum::{ApiRouter, routing::get_with}, + transform::TransformOperation, +}; use axum::{ Extension, Json, extract::{Path, Query, State}, http::StatusCode, - response::{IntoResponse, Redirect}, + response::Redirect, }; +use schemars::JsonSchema; use serde::Deserialize; -use utoipa::{IntoParams, ToSchema}; -use utoipa_axum::{router::OpenApiRouter, routes}; use crate::{ api::{ApiTag, RoutePrefix}, db::models::ChatRsUser, error::AppResult, - extractors::{CurrentUser, Database, PublicAuthConfig, SessionMeta}, + extractors::{AppSession, CurrentUser, Database, PublicAuthConfig}, services::auth::oauth::OAuthProviderEnum, state::AppState, }; -pub fn routes() -> OpenApiRouter { - OpenApiRouter::new() - .routes(routes!(get_user)) - .routes(routes!(get_config)) - .routes(routes!(oauth_login)) - .routes(routes!(oauth_login_callback)) - .routes(routes!(logout)) +pub fn routes() -> ApiRouter { + ApiRouter::new() + .api_route("/user", get_with(get_user, get_user_docs)) + .api_route("/config", get_with(get_config, get_config_docs)) + .api_route("/login", get_with(oauth_login, oauth_login_docs)) + .api_route( + "/login/callback", + get_with(oauth_callback, oauth_callback_docs), + ) + .api_route( + "/logout", + get_with(logout, logout_docs).post_with(logout, logout_docs), + ) + .with_path_items(|op| op.tag(ApiTag::Auth.into())) } -/// Get current user -#[utoipa::path( - get, path = "/user", - responses((status = OK, body = ChatRsUser)), - tag = ApiTag::Auth.into()) -] +fn get_user_docs(op: TransformOperation) -> TransformOperation { + op.id("get_user").summary("Get current user") +} async fn get_user( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, State(state): State, -) -> AppResult { +) -> AppResult> { let user = state.auth_service().get_user(&mut db, &user_id).await?; Ok(Json(user)) } -#[utoipa::path( - get, path = "/config", - responses((status = OK, body = PublicAuthConfig)), - tag = ApiTag::Auth.into() -)] -async fn get_config(auth_config: PublicAuthConfig) -> impl IntoResponse { +fn get_config_docs(op: TransformOperation) -> TransformOperation { + op.id("get_auth_config").summary("Get auth config") +} +async fn get_config(auth_config: PublicAuthConfig) -> Json { Json(auth_config) } -fn oauth_callback_path(route_prefix: &'static str, provider: OAuthProviderEnum) -> String { +fn oauth_callback_path(route_prefix: &'static str, provider: &OAuthProviderEnum) -> String { format!("{route_prefix}/login/{provider}/callback") } -/// OAuth login redirect -#[utoipa::path( - get, path = "/login/{provider}", - params(("provider" = OAuthProviderEnum, Path)), - responses((status = OK)), - tag = ApiTag::Auth.into(), -)] +fn oauth_login_docs(op: TransformOperation) -> TransformOperation { + op.id("oauth_login") + .summary("OAuth login") + .description("OAuth login redirect") +} async fn oauth_login( Path(provider): Path, Extension(RoutePrefix(prefix)): Extension, State(state): State, - session: tower_sessions::Session, -) -> AppResult { + AppSession { session, .. }: AppSession, +) -> AppResult { let oauth = state.auth_service().oauth(); let auth_url = oauth - .authorize_url(provider, &oauth_callback_path(prefix, provider), &session) + .authorize_url(&provider, &oauth_callback_path(prefix, &provider), &session) .await?; Ok(Redirect::to(auth_url.as_str())) } -#[derive(Debug, Clone, Deserialize, IntoParams, ToSchema)] +#[derive(Debug, Deserialize, JsonSchema)] struct OAuthCallbackQuery { code: String, state: String, } -/// OAuth login callback -#[utoipa::path( - get, path = "/login/{provider}/callback", - params( - ("query" = inline(OAuthCallbackQuery), Query), - ("provider" = OAuthProviderEnum, Path, description = "the OAuth provider") - ), - responses((status = OK)), - tag = ApiTag::Auth.into(), -)] -async fn oauth_login_callback( +fn oauth_callback_docs(op: TransformOperation) -> TransformOperation { + op.id("oauth_login_callback") + .summary("OAuth login callback") +} +async fn oauth_callback( Path(provider): Path, Query(query): Query, + maybe_user: Option, Extension(RoutePrefix(prefix)): Extension, + AppSession { session, meta }: AppSession, Database(mut db): Database, - State(state): State, - session: tower_sessions::Session, - meta: SessionMeta, - maybe_user: Option, -) -> AppResult { - let oauth = state.auth_service().oauth(); + State(app_state): State, +) -> AppResult { + let oauth = app_state.auth_service().oauth(); let token = oauth .exchange_code( - provider, - &oauth_callback_path(prefix, provider), + &provider, + &oauth_callback_path(prefix, &provider), &session, &query.code, &query.state, ) .await?; let user = oauth - .get_user(&mut db, provider, &token, maybe_user) + .get_user(&mut db, &provider, &token, maybe_user) .await?; - state + app_state .auth_service() .session() .login(&session, &meta, &user.id) .await?; - Ok(Redirect::to(&state.config.server.base_url)) + Ok(Redirect::to(&app_state.config.server.base_url)) } -/// Logout -#[utoipa::path( - method(get, post), path = "/logout", - tag = ApiTag::Auth.into(), - responses((status = NO_CONTENT)), -)] +fn logout_docs(op: TransformOperation) -> TransformOperation { + op.id("logout").summary("Logout") +} async fn logout( - session: tower_sessions::Session, + AppSession { session, .. }: AppSession, State(state): State, -) -> AppResult { +) -> AppResult { state.auth_service().session().logout(&session).await?; Ok(StatusCode::NO_CONTENT) } diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index 847cdf1..fe557e6 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -1,26 +1,74 @@ -use axum::{ - Json, - extract::{Path, State}, - response::IntoResponse, -}; +use aide::axum::ApiRouter; +use axum::{Json, extract::State}; +use axum_typed_routing::{TypedApiRouter, api_route}; +use derive_more::Into; +use schemars::JsonSchema; use serde::{Deserialize, Serialize}; -use utoipa::ToSchema; -use utoipa_axum::{router::OpenApiRouter, routes}; use uuid::Uuid; use crate::{ api::ApiTag, - error::AppError, + error::{AppError, AppResult}, extractors::{CurrentUser, Database}, llm::types::{LlmChatOptions, LlmUserMessage}, state::AppState, }; -pub fn routes() -> OpenApiRouter { - OpenApiRouter::new().routes(routes!(chat_stream)) +pub fn routes() -> ApiRouter { + ApiRouter::new() + .typed_api_route(prompt) + .typed_api_route(chat_stream) +} + +#[derive(Debug, Deserialize, JsonSchema)] +struct PromptInput { + /// The prompt to send to the LLM provider + message: String, + /// The ID of the provider to chat with + provider_id: i32, + /// Configuration for the provider + options: LlmChatOptions, } -#[derive(Debug, Deserialize)] +#[api_route(POST "/prompt" { + summary: "Prompt", + description: "Send a simple prompt to a provider and get the response", + transform: |op| op.tag(ApiTag::Chat.into()), +})] +async fn prompt( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, + State(state): State, + Json(input): Json, +) -> AppResult> { + let llm_provider = state + .provider_service() + .build_llm_provider(&mut db, &user_id, input.provider_id) + .await?; + let text = state + .chat_service() + .prompt( + llm_provider, + LlmUserMessage { + text: input.message, + ..Default::default() + }, + input.options, + ) + .await?; + + Ok(Json(PromptResponse { text })) +} + +#[derive(Serialize, JsonSchema)] +struct PromptResponse { + text: String, +} + +#[derive(Into, Deserialize, JsonSchema)] +struct SessionIdPath(Uuid); + +#[derive(Debug, Deserialize, JsonSchema)] struct ChatInput { /// The new chat message from the user message: Option, @@ -30,22 +78,18 @@ struct ChatInput { options: LlmChatOptions, } -/// Streaming chat -/// -/// Send a message in a chat session and stream the response -#[utoipa::path( - get, path = "/{session_id}", - params(("session_id" = Uuid, Path)), - responses((status = OK, body = StreamAccess)), - tag = ApiTag::Chat.into(), -)] +#[api_route(POST "/{session_id}" { + summary: "Streaming chat", + description: "Send a message in a chat session and stream the response", + transform: |op| op.tag(ApiTag::Chat.into()), +})] async fn chat_stream( + session_id: SessionIdPath, CurrentUser { user_id }: CurrentUser, - Path(session_id): Path, Database(mut db): Database, State(state): State, Json(input): Json, -) -> Result { +) -> Result, AppError> { let llm_provider = state .provider_service() .build_llm_provider(&mut db, &user_id, input.provider_id) @@ -55,7 +99,7 @@ async fn chat_stream( .stream_user_chat( &mut db, user_id, - session_id, + session_id.into(), input.provider_id, llm_provider, input @@ -71,7 +115,8 @@ async fn chat_stream( })) } -#[derive(Serialize, ToSchema)] +/// The URL and Bearer token to access the SSE stream +#[derive(Serialize, JsonSchema)] struct StreamAccess { url: String, token: String, diff --git a/server-new/src/api/health.rs b/server-new/src/api/health.rs index 8f09846..65671ea 100644 --- a/server-new/src/api/health.rs +++ b/server-new/src/api/health.rs @@ -1,12 +1,13 @@ -use utoipa_axum::{router::OpenApiRouter, routes}; +use aide::axum::ApiRouter; +use axum_typed_routing::{TypedApiRouter, api_route}; use crate::state::AppState; -pub fn routes() -> OpenApiRouter { - OpenApiRouter::new().routes(routes!(health_handler)) +pub fn routes() -> ApiRouter { + ApiRouter::new().typed_api_route(health) } -#[utoipa::path(get, path = "", responses((status = OK, body = &str)))] -async fn health_handler() -> &'static str { +#[api_route(GET "/" with AppState { summary: "Health route" })] +async fn health() -> &'static str { "OK" } diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index ccb1bfb..3803fc3 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -1,43 +1,40 @@ -use axum::Extension; +use std::sync::Arc; + +use aide::{axum::ApiRouter, openapi::OpenApi, swagger::Swagger}; +use axum::{Extension, http::header, response::IntoResponse, routing::get}; use axum_plugin::AdHocPlugin; -use strum::{AsRefStr, IntoStaticStr}; -use utoipa::OpenApi; -use utoipa_axum::router::OpenApiRouter; -use utoipa_scalar::{Scalar, Servable}; +use strum::{AsRefStr, Display, EnumIter, EnumMessage, IntoEnumIterator, IntoStaticStr}; -use crate::{services::auth::oauth::OAuthProviderEnum, state::AppState}; +use crate::state::AppState; pub mod api_key; pub mod auth; pub mod chat; pub mod health; -#[derive(AsRefStr, IntoStaticStr)] -#[strum(serialize_all = "snake_case")] +#[derive(Display, AsRefStr, IntoStaticStr, EnumMessage, EnumIter)] enum ApiTag { + #[strum(message = "Manage API keys")] ApiKey, + #[strum(message = "Authentication")] Auth, + #[strum(message = "Chats and sessions")] Chat, } -#[derive(OpenApi)] -#[openapi( - servers((url = "/api")), - components( - schemas(OAuthProviderEnum) - ), - tags( - (name = ApiTag::ApiKey.as_ref(), description = "Manage API keys"), - (name = ApiTag::Auth.as_ref(), description = "Authentication"), - (name = ApiTag::Chat.as_ref(), description = "Chats and sessions") - ) -)] -struct ApiDoc; - /// Adds all API routes with OpenAPI docs to the server under `/api` pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("API routes").on_setup(|router, _state| { - let (api_routes, openapi) = OpenApiRouter::with_openapi(ApiDoc::openapi()) + let mut openapi = OpenApi::default(); + for tag in ApiTag::iter() { + openapi.tags.push(aide::openapi::Tag { + name: tag.to_string(), + description: tag.get_message().map(String::from), + ..Default::default() + }) + } + + let api_routes = ApiRouter::new() .nest("/api_key", api_key::routes()) .nest( "/auth", @@ -45,11 +42,28 @@ pub fn plugin() -> AdHocPlugin { ) .nest("/chat", chat::routes()) .nest("/health", health::routes()) - .split_for_parts(); + .finish_api(&mut openapi); - Ok(router.nest("/api", api_routes.merge(Scalar::with_url("/docs", openapi)))) + let api_routes_with_docs = api_routes + .route( + "/docs/openapi.json", + get(openapi_route).layer(Extension(Arc::new(openapi))), + ) + .route( + "/docs", + get(Swagger::new("/api/docs/openapi.json") + .with_title("RsChat API") + .axum_handler()), + ); + + Ok(router.nest("/api", api_routes_with_docs)) }) } +async fn openapi_route(Extension(openapi): Extension>) -> impl IntoResponse { + axum::Json(openapi) +} + +/// Extension to pass the route prefix to child routes #[derive(Clone)] struct RoutePrefix(&'static str); diff --git a/server-new/src/db/models/api_key.rs b/server-new/src/db/models/api_key.rs index 2cc1363..ad9213a 100644 --- a/server-new/src/db/models/api_key.rs +++ b/server-new/src/db/models/api_key.rs @@ -1,12 +1,12 @@ use chrono::{DateTime, Utc}; use diesel::prelude::*; +use schemars::JsonSchema; use serde::Serialize; -use utoipa::ToSchema; use uuid::Uuid; use crate::db::models::ChatRsUser; -#[derive(Identifiable, Queryable, Selectable, Associations, Serialize, ToSchema)] +#[derive(Identifiable, Queryable, Selectable, Associations, Serialize, JsonSchema)] #[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] #[diesel(table_name = super::schema::app_api_keys)] pub struct ChatRsApiKey { diff --git a/server-new/src/db/models/user.rs b/server-new/src/db/models/user.rs index fe5842e..bb47d5d 100644 --- a/server-new/src/db/models/user.rs +++ b/server-new/src/db/models/user.rs @@ -1,11 +1,11 @@ use diesel::prelude::*; +use schemars::JsonSchema; use serde::Serialize; use serde_with::skip_serializing_none; -use utoipa::ToSchema; use uuid::Uuid; #[skip_serializing_none] -#[derive(Identifiable, Queryable, Selectable, Serialize, ToSchema)] +#[derive(Identifiable, Queryable, Selectable, Serialize, JsonSchema)] #[diesel(table_name = super::schema::users)] pub struct ChatRsUser { pub id: Uuid, diff --git a/server-new/src/error.rs b/server-new/src/error.rs index a9dce07..ef4e324 100644 --- a/server-new/src/error.rs +++ b/server-new/src/error.rs @@ -1,3 +1,4 @@ +use aide::OperationIo; use axum::{ Json, http::StatusCode, @@ -7,11 +8,11 @@ use serde::Serialize; use crate::db::DbPoolError; -/// Global result type that can be used for API route handlers +/// Global API result type that can be used in route handlers pub type AppResult = Result; -/// Global error type -#[derive(Debug)] +/// Global API error type +#[derive(Debug, OperationIo)] pub struct AppError { status: StatusCode, message: String, diff --git a/server-new/src/extractors/auth_config.rs b/server-new/src/extractors/auth_config.rs index f0adb49..2d93571 100644 --- a/server-new/src/extractors/auth_config.rs +++ b/server-new/src/extractors/auth_config.rs @@ -1,13 +1,14 @@ +use aide::OperationIo; use axum::extract::FromRequestParts; +use schemars::JsonSchema; use serde::Serialize; use serde_with::skip_serializing_none; -use utoipa::ToSchema; use crate::state::AppState; /// The current auth configuration of the server #[skip_serializing_none] -#[derive(Debug, Serialize, ToSchema)] +#[derive(Debug, Serialize, JsonSchema, OperationIo)] pub struct PublicAuthConfig { /// Whether GitHub login is enabled github: bool, @@ -21,7 +22,7 @@ pub struct PublicAuthConfig { // sso: Option, } -#[derive(Debug, Serialize, ToSchema)] +#[derive(Debug, Serialize, JsonSchema)] struct Oidc { /// The name of the OIDC provider name: String, diff --git a/server-new/src/extractors/database.rs b/server-new/src/extractors/database.rs index fccfda1..6faa586 100644 --- a/server-new/src/extractors/database.rs +++ b/server-new/src/extractors/database.rs @@ -1,8 +1,10 @@ +use aide::OperationIo; use axum::extract::FromRequestParts; use crate::{db::DbService, error::AppError, state::AppState}; /// An extractor to retrieve a database connection from the pool +#[derive(OperationIo)] pub struct Database(pub DbService); impl FromRequestParts for Database { diff --git a/server-new/src/extractors/mod.rs b/server-new/src/extractors/mod.rs index 1ffb082..2871247 100644 --- a/server-new/src/extractors/mod.rs +++ b/server-new/src/extractors/mod.rs @@ -7,5 +7,5 @@ mod user; pub use auth_config::PublicAuthConfig; pub use database::Database; -pub use session::SessionMeta; +pub use session::{AppSession, SessionMeta}; pub use user::CurrentUser; diff --git a/server-new/src/extractors/session.rs b/server-new/src/extractors/session.rs index 368207a..1e5c372 100644 --- a/server-new/src/extractors/session.rs +++ b/server-new/src/extractors/session.rs @@ -3,24 +3,52 @@ use std::{ str::FromStr, }; +use aide::OperationIo; +use anyhow::anyhow; use axum::{ extract::{ConnectInfo, FromRequestParts}, http::header, }; use serde::{Deserialize, Serialize}; +use tower_sessions::Session; -use crate::{db::UtcDateTime, state::AppState}; +use crate::{db::UtcDateTime, error::AppError, state::AppState}; + +/// Extractor to get raw session and request metadata +#[derive(OperationIo)] +pub struct AppSession { + pub session: Session, + pub meta: SessionMeta, +} /// Session metadata extracted on login. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, OperationIo)] pub struct SessionMeta { pub start_time: UtcDateTime, pub ip: Option, pub user_agent: Option, } +impl FromRequestParts for AppSession { + type Rejection = AppError; + + async fn from_request_parts( + parts: &mut axum::http::request::Parts, + state: &AppState, + ) -> Result { + let meta = SessionMeta::from_request_parts(parts, state).await?; + let session = parts + .extensions + .get::() + .cloned() + .ok_or_else(|| AppError::internal(anyhow!("session not attached to request")))?; + + Ok(Self { session, meta }) + } +} + impl FromRequestParts for SessionMeta { - type Rejection = (); + type Rejection = AppError; async fn from_request_parts( parts: &mut axum::http::request::Parts, diff --git a/server-new/src/extractors/user.rs b/server-new/src/extractors/user.rs index 427cff0..ea9e473 100644 --- a/server-new/src/extractors/user.rs +++ b/server-new/src/extractors/user.rs @@ -1,3 +1,4 @@ +use aide::OperationIo; use anyhow::anyhow; use axum::{ extract::{FromRequestParts, OptionalFromRequestParts}, @@ -16,7 +17,7 @@ if there is no active user. - If used as `Option`, will be `Some` if there is an active user and `None` otherwise. */ -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, OperationIo)] pub struct CurrentUser { pub user_id: Uuid, } diff --git a/server-new/src/llm/types.rs b/server-new/src/llm/types.rs index 5f6e8fc..be697d9 100644 --- a/server-new/src/llm/types.rs +++ b/server-new/src/llm/types.rs @@ -1,3 +1,4 @@ +use schemars::JsonSchema; use serde::{Deserialize, Serialize}; /// Generic LLM prompt @@ -22,7 +23,7 @@ pub enum LlmMessage { } /// Generic chat options for all LLM providers -#[derive(Clone, Debug, Default, Serialize, Deserialize)] +#[derive(Clone, Debug, Default, Serialize, Deserialize, JsonSchema)] pub struct LlmChatOptions { pub model: String, pub temperature: Option, diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index fe0fdc9..c26a08d 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -1,6 +1,7 @@ use std::collections::HashMap; use futures::future::BoxFuture; +use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use simple_oauth::{ SimpleOAuthClient, SimpleOAuthError, SimpleOAuthProvider, @@ -8,7 +9,6 @@ use simple_oauth::{ }; use strum::Display; use tower_sessions::Session; -use utoipa::ToSchema; use crate::{ config::AppConfig, @@ -36,7 +36,7 @@ pub type OAuthProviderMap = HashMap>; /// Supported OAuth provider -#[derive(Debug, Display, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, ToSchema)] +#[derive(Debug, Display, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, JsonSchema)] #[serde(rename_all = "lowercase")] #[strum(serialize_all = "lowercase")] pub enum OAuthProviderEnum { @@ -79,11 +79,11 @@ impl<'a> OAuthService<'a> { fn oauth_provider( &self, - provider: OAuthProviderEnum, + provider: &OAuthProviderEnum, ) -> AuthResult<(&OAuthClient, &Box)> { let (client, provider) = self .provider_map - .get(&provider) + .get(provider) .ok_or_else(|| AuthError::BadRequest("unsupported OAuth provider"))?; Ok((client, provider)) } @@ -94,7 +94,7 @@ impl<'a> OAuthService<'a> { pub async fn authorize_url( &self, - provider: OAuthProviderEnum, + provider: &OAuthProviderEnum, callback_path: &str, session: &Session, ) -> AuthResult { @@ -114,7 +114,7 @@ impl<'a> OAuthService<'a> { pub async fn exchange_code( &self, - provider: OAuthProviderEnum, + provider: &OAuthProviderEnum, callback_path: &str, session: &Session, code: &str, @@ -148,7 +148,7 @@ impl<'a> OAuthService<'a> { pub async fn get_user( &self, db: &mut DbService, - provider: OAuthProviderEnum, + provider: &OAuthProviderEnum, token: &StandardTokenResponse, active_session: Option, ) -> AuthResult { diff --git a/server-new/src/services/auth/session.rs b/server-new/src/services/auth/session.rs index 6bc93ee..039500b 100644 --- a/server-new/src/services/auth/session.rs +++ b/server-new/src/services/auth/session.rs @@ -51,6 +51,16 @@ impl AuthSessionService { Ok(user_id) } + /// Extract the user id from the raw session hashmap + pub(super) fn user_id_from_record_data( + data: &HashMap, + ) -> StoreResult> { + data.get(USER_ID_FIELD) + .and_then(|val| val.as_str().map(Uuid::try_parse)) + .transpose() + .map_err(|_| StoreError::Decode("invalid user id field".into())) + } + /// Logout the user, deleting the current session. pub async fn logout(&self, session: &Session) -> AuthResult<()> { Ok(session.flush().await?) @@ -63,12 +73,3 @@ impl AuthSessionService { Ok(db.auth_sessions().delete_expired().await?) } } - -pub(super) fn user_id_from_record_data( - data: &HashMap, -) -> StoreResult> { - data.get(USER_ID_FIELD) - .and_then(|val| val.as_str().map(Uuid::try_parse)) - .transpose() - .map_err(|_| StoreError::Encode("invalid user id field".to_owned())) -} diff --git a/server-new/src/services/auth/session_store.rs b/server-new/src/services/auth/session_store.rs index 8237e1d..24e97d8 100644 --- a/server-new/src/services/auth/session_store.rs +++ b/server-new/src/services/auth/session_store.rs @@ -7,7 +7,10 @@ use tower_sessions::{ }; use uuid::Uuid; -use crate::db::{DbPool, DbService, UtcDateTime}; +use crate::{ + db::{DbPool, DbService, UtcDateTime}, + services::auth::session::AuthSessionService, +}; #[derive(Clone)] pub struct SessionDbStore { @@ -46,7 +49,7 @@ impl SessionStore for SessionDbStore { /// Creates a new session in the store with the provided session record. async fn create(&self, record: &mut Record) -> Result<()> { let session_id = Self::get_session_uuid(&record.id); - let user_id = super::session::user_id_from_record_data(&record.data)?; + let user_id = AuthSessionService::user_id_from_record_data(&record.data)?; let expires_at = Self::convert_expiry(record.expiry_date)?; let mut db = self.get_db().await?; diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index b3ce9e5..e68208b 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -13,7 +13,7 @@ use crate::{ }, llm::{ interface::LlmProvider, - types::{LlmChatOptions, LlmChatRequest, LlmUserMessage}, + types::{LlmChatOptions, LlmChatRequest, LlmPrompt, LlmUserMessage}, }, services::{ chat::error::ChatError, @@ -40,6 +40,20 @@ impl<'r> ChatService<'r> { } } + /// Send a simple prompt to the LLM provider + pub async fn prompt( + &self, + provider: Arc, + prompt: LlmUserMessage, + options: LlmChatOptions, + ) -> Result { + let llm_prompt = LlmPrompt { + text: &prompt.text, + options: &options, + }; + Ok(provider.prompt(llm_prompt).await?) + } + pub async fn stream_user_chat( &self, db: &mut DbService, @@ -107,7 +121,7 @@ impl<'r> ChatService<'r> { let response = StreamingService::process_stream(stream, ws_writer, ws_reader).await; let stream_cancelled = response.cancelled; if let Err(err) = - Self::persist_response(db_pool, &session_id, provider_id, chat_options, response) + Self::persist_response(db_pool, session_id, provider_id, chat_options, response) .await { tracing::error!("Failed to save assistant response: {err}"); @@ -127,7 +141,7 @@ impl<'r> ChatService<'r> { /// Save response message and metadata to database async fn persist_response( db_pool: DbPool, - session_id: &Uuid, + session_id: Uuid, provider_id: i32, chat_options: LlmChatOptions, response: LlmStreamOutput, From f398c25decd1814bcf7e97b2759bf9c49dfda794 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sun, 5 Jul 2026 22:26:01 -0400 Subject: [PATCH 062/111] api doc - checkpoint 2 --- server-new/Cargo.lock | 10 ++ server-new/Cargo.toml | 1 + server-new/crates/aide-docs-macro/Cargo.toml | 14 +++ server-new/crates/aide-docs-macro/src/lib.rs | 64 +++++++++++++ server-new/src/api/api_key.rs | 96 ++++++-------------- server-new/src/api/auth.rs | 29 ++---- server-new/src/api/chat.rs | 40 ++++---- server-new/src/db/repositories/api_key.rs | 13 ++- 8 files changed, 149 insertions(+), 118 deletions(-) create mode 100644 server-new/crates/aide-docs-macro/Cargo.toml create mode 100644 server-new/crates/aide-docs-macro/src/lib.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index f6ebf95..841b197 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -123,6 +123,15 @@ dependencies = [ "tracing", ] +[[package]] +name = "aide-docs-macro" +version = "0.1.0" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "aide-macros" version = "0.16.0-alpha.4" @@ -2395,6 +2404,7 @@ version = "0.1.0" dependencies = [ "aes-gcm 0.11.0", "aide 0.16.0-alpha.4", + "aide-docs-macro", "anyhow", "async-stream", "async-trait", diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 3db8c8a..7738a41 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -12,6 +12,7 @@ aide = { rev = "7246c20", features = ["axum", "axum-json", "axum-query", "macros", "swagger"] } +aide-docs-macro = { path = "crates/aide-docs-macro" } anyhow = "1.0.102" async-stream = "0.3.6" async-trait = "0.1.89" diff --git a/server-new/crates/aide-docs-macro/Cargo.toml b/server-new/crates/aide-docs-macro/Cargo.toml new file mode 100644 index 0000000..b655ebc --- /dev/null +++ b/server-new/crates/aide-docs-macro/Cargo.toml @@ -0,0 +1,14 @@ +[package] +name = "aide-docs-macro" +version = "0.1.0" +edition = "2024" +description = "Internal attribute macro for API operation docs" +publish = false + +[lib] +proc-macro = true + +[dependencies] +proc-macro2 = "1" +quote = "1" +syn = { version = "2", features = ["full", "parsing"] } diff --git a/server-new/crates/aide-docs-macro/src/lib.rs b/server-new/crates/aide-docs-macro/src/lib.rs new file mode 100644 index 0000000..5c46194 --- /dev/null +++ b/server-new/crates/aide-docs-macro/src/lib.rs @@ -0,0 +1,64 @@ +use proc_macro::TokenStream; +use quote::{format_ident, quote}; +use syn::{ItemFn, LitStr, Result, Token, parse::Parse, parse::ParseStream, parse_macro_input}; + +struct DocsArgs { + summary: LitStr, + description: Option, +} + +impl Parse for DocsArgs { + fn parse(input: ParseStream<'_>) -> Result { + let summary = input.parse()?; + let description = if input.peek(Token![,]) { + input.parse::()?; + Some(input.parse()?) + } else { + None + }; + + if !input.is_empty() { + input.parse::()?; + } + + Ok(Self { + summary, + description, + }) + } +} + +/// Convenience macro for generating API docs with aide. Generates a function +/// called `_docs` that can be passed as the transform function +/// to `get_with`, `post_with`, etc. +/// +/// # Syntax +/// `#[docs("", ""]` +#[proc_macro_attribute] +pub fn docs(args: TokenStream, input: TokenStream) -> TokenStream { + let args = parse_macro_input!(args as DocsArgs); + let handler = parse_macro_input!(input as ItemFn); + let handler_name = &handler.sig.ident; + let docs_name = format_ident!("{}_docs", handler_name); + let operation_id = handler_name.to_string(); + let summary = args.summary; + + let description = args.description.map(|description| { + quote! { + .description(#description) + } + }); + + quote! { + fn #docs_name( + op: ::aide::transform::TransformOperation, + ) -> ::aide::transform::TransformOperation { + op.id(#operation_id) + .summary(#summary) + #description + } + + #handler + } + .into() +} diff --git a/server-new/src/api/api_key.rs b/server-new/src/api/api_key.rs index 61ed15e..252b4c8 100644 --- a/server-new/src/api/api_key.rs +++ b/server-new/src/api/api_key.rs @@ -1,7 +1,13 @@ -use aide::{OperationInput, OperationIo, axum::ApiRouter}; -use axum::{Json, extract::State, http::StatusCode}; -use axum_typed_routing::{TypedApiRouter, api_route}; -use derive_more::Deref; +use aide::axum::{ + ApiRouter, + routing::{delete_with, get_with, post_with}, +}; +use aide_docs_macro::docs; +use axum::{ + Json, + extract::{Path, State}, + http::StatusCode, +}; use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use uuid::Uuid; @@ -9,22 +15,20 @@ use uuid::Uuid; use crate::{ api::ApiTag, db::models::ChatRsApiKey, - error::AppResult, + error::{AppError, AppResult}, extractors::{CurrentUser, Database}, state::AppState, }; pub fn routes() -> ApiRouter { ApiRouter::new() - .typed_api_route(list_api_keys) - .typed_api_route(create_api_key) - .typed_api_route(delete_api_key) + .api_route("/", get_with(list_api_keys, list_api_keys_docs)) + .api_route("/", post_with(create_api_key, create_api_key_docs)) + .api_route("/{id}", delete_with(delete_api_key, delete_api_key_docs)) + .with_path_items(|op| op.tag(ApiTag::ApiKey.into())) } -#[api_route(GET "/" with AppState { - summary: "List API keys", - transform: |op| op.tag(ApiTag::ApiKey.into()), -})] +#[docs("List API keys")] async fn list_api_keys( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, @@ -38,16 +42,7 @@ struct ApiKeyCreateInput { name: String, } -#[derive(Serialize, JsonSchema)] -struct ApiKeyCreateResponse { - id: Uuid, - key: String, -} - -#[api_route(POST "/" { - summary: "Create API key", - transform: |op| op.tag(ApiTag::ApiKey.into()), -})] +#[docs("Create API key")] async fn create_api_key( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, @@ -59,60 +54,23 @@ async fn create_api_key( .api_keys() .create_api_key(&mut db, &user_id, &input.name) .await?; - Ok(Json(ApiKeyCreateResponse { id, key })) } -#[derive(Deref, Deserialize, OperationIo, JsonSchema)] -pub struct ApiKeyPath { - pub id: Uuid, -} - -#[derive(Deref, Serialize, Deserialize, JsonSchema)] -pub struct UuidPath(pub Uuid); -impl OperationInput for UuidPath { - fn operation_input( - ctx: &mut aide::generate::GenContext, - operation: &mut aide::openapi::Operation, - ) { - use aide::openapi::{ - Parameter, ParameterData, ParameterSchemaOrContent, ReferenceOr, SchemaObject, - }; - - operation - .parameters - .push(ReferenceOr::Item(Parameter::Path { - parameter_data: ParameterData { - name: UuidPath::schema_name().into(), - description: None, - required: true, - deprecated: Default::default(), - format: ParameterSchemaOrContent::Schema(SchemaObject { - json_schema: UuidPath::json_schema(&mut ctx.schema), - external_docs: None, - example: None, - }), - example: Default::default(), - examples: Default::default(), - explode: Default::default(), - extensions: Default::default(), - }, - style: aide::openapi::PathStyle::Simple, - })) - } +#[derive(Serialize, JsonSchema)] +struct ApiKeyCreateResponse { + id: Uuid, + key: String, } -#[api_route(DELETE "/{id}" with AppState { - summary: "Delete API key", - responses: { 204: () }, - transform: |op| op.tag(ApiTag::ApiKey.into()), -})] +#[docs("Delete API key")] async fn delete_api_key( - id: UuidPath, + Path(id): Path, CurrentUser { user_id }: CurrentUser, Database(mut db): Database, ) -> AppResult { - let _ = db.api_keys().delete(&user_id, &id).await?; - - Ok(StatusCode::NO_CONTENT) + match db.api_keys().delete(&user_id, &id).await? { + Some(_) => Ok(StatusCode::NO_CONTENT), + None => Err(AppError::not_found("API key not found")), + } } diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index 63e2ec9..23f836e 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -1,7 +1,5 @@ -use aide::{ - axum::{ApiRouter, routing::get_with}, - transform::TransformOperation, -}; +use aide::axum::{ApiRouter, routing::get_with}; +use aide_docs_macro::docs; use axum::{ Extension, Json, extract::{Path, Query, State}, @@ -36,9 +34,7 @@ pub fn routes() -> ApiRouter { .with_path_items(|op| op.tag(ApiTag::Auth.into())) } -fn get_user_docs(op: TransformOperation) -> TransformOperation { - op.id("get_user").summary("Get current user") -} +#[docs("Get user", "Get the current user")] async fn get_user( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, @@ -48,9 +44,7 @@ async fn get_user( Ok(Json(user)) } -fn get_config_docs(op: TransformOperation) -> TransformOperation { - op.id("get_auth_config").summary("Get auth config") -} +#[docs("Get auth config", "Get the current auth configuration of the server")] async fn get_config(auth_config: PublicAuthConfig) -> Json { Json(auth_config) } @@ -59,11 +53,7 @@ fn oauth_callback_path(route_prefix: &'static str, provider: &OAuthProviderEnum) format!("{route_prefix}/login/{provider}/callback") } -fn oauth_login_docs(op: TransformOperation) -> TransformOperation { - op.id("oauth_login") - .summary("OAuth login") - .description("OAuth login redirect") -} +#[docs("OAuth login", "OAuth login redirect")] async fn oauth_login( Path(provider): Path, Extension(RoutePrefix(prefix)): Extension, @@ -84,10 +74,7 @@ struct OAuthCallbackQuery { state: String, } -fn oauth_callback_docs(op: TransformOperation) -> TransformOperation { - op.id("oauth_login_callback") - .summary("OAuth login callback") -} +#[docs("OAuth login callback sdf")] async fn oauth_callback( Path(provider): Path, Query(query): Query, @@ -119,9 +106,7 @@ async fn oauth_callback( Ok(Redirect::to(&app_state.config.server.base_url)) } -fn logout_docs(op: TransformOperation) -> TransformOperation { - op.id("logout").summary("Logout") -} +#[docs("Logout")] async fn logout( AppSession { session, .. }: AppSession, State(state): State, diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index fe557e6..eabcc8c 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -1,13 +1,14 @@ -use aide::axum::ApiRouter; -use axum::{Json, extract::State}; -use axum_typed_routing::{TypedApiRouter, api_route}; -use derive_more::Into; +use aide::axum::{ApiRouter, routing::post_with}; +use aide_docs_macro::docs; +use axum::{ + Json, + extract::{Path, State}, +}; use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use uuid::Uuid; use crate::{ - api::ApiTag, error::{AppError, AppResult}, extractors::{CurrentUser, Database}, llm::types::{LlmChatOptions, LlmUserMessage}, @@ -16,8 +17,11 @@ use crate::{ pub fn routes() -> ApiRouter { ApiRouter::new() - .typed_api_route(prompt) - .typed_api_route(chat_stream) + .api_route("/prompt", post_with(prompt, prompt_docs)) + .api_route( + "/session/{session_id}", + post_with(chat_stream, chat_stream_docs), + ) } #[derive(Debug, Deserialize, JsonSchema)] @@ -30,11 +34,7 @@ struct PromptInput { options: LlmChatOptions, } -#[api_route(POST "/prompt" { - summary: "Prompt", - description: "Send a simple prompt to a provider and get the response", - transform: |op| op.tag(ApiTag::Chat.into()), -})] +#[docs("Prompt", "Send a single prompt to a provider and get the response")] async fn prompt( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, @@ -65,9 +65,6 @@ struct PromptResponse { text: String, } -#[derive(Into, Deserialize, JsonSchema)] -struct SessionIdPath(Uuid); - #[derive(Debug, Deserialize, JsonSchema)] struct ChatInput { /// The new chat message from the user @@ -78,13 +75,12 @@ struct ChatInput { options: LlmChatOptions, } -#[api_route(POST "/{session_id}" { - summary: "Streaming chat", - description: "Send a message in a chat session and stream the response", - transform: |op| op.tag(ApiTag::Chat.into()), -})] +#[docs( + "Send chat", + "Send a message in a chat session and stream the response" +)] async fn chat_stream( - session_id: SessionIdPath, + Path(session_id): Path, CurrentUser { user_id }: CurrentUser, Database(mut db): Database, State(state): State, @@ -99,7 +95,7 @@ async fn chat_stream( .stream_user_chat( &mut db, user_id, - session_id.into(), + session_id, input.provider_id, llm_provider, input diff --git a/server-new/src/db/repositories/api_key.rs b/server-new/src/db/repositories/api_key.rs index 42aed9a..ac255e9 100644 --- a/server-new/src/db/repositories/api_key.rs +++ b/server-new/src/db/repositories/api_key.rs @@ -47,15 +47,18 @@ impl<'a> ApiKeyRepository<'a> { Ok(id) } - pub async fn delete(&mut self, user_id: &Uuid, api_key_id: &Uuid) -> Result { - let id: Uuid = diesel::delete(app_api_keys::table) + pub async fn delete( + &mut self, + user_id: &Uuid, + api_key_id: &Uuid, + ) -> Result, Error> { + diesel::delete(app_api_keys::table) .filter(app_api_keys::id.eq(api_key_id)) .filter(app_api_keys::user_id.eq(user_id)) .returning(app_api_keys::id) .get_result(self.db) - .await?; - - Ok(id) + .await + .optional() } pub async fn delete_by_user(&mut self, user_id: &Uuid) -> Result, Error> { From 1506b7fbee5c00253be54c967e56c7ab40c73f77 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sun, 5 Jul 2026 23:18:57 -0400 Subject: [PATCH 063/111] api doc completed --- server-new/Cargo.lock | 166 +--------------- server-new/Cargo.toml | 2 - server-new/crates/aide-docs-macro/src/lib.rs | 195 ++++++++++++++++++- server-new/src/api/api_key.rs | 21 +- server-new/src/api/auth.rs | 30 +-- server-new/src/api/chat.rs | 23 +-- server-new/src/api/health.rs | 8 +- server-new/src/api/mod.rs | 42 ++-- 8 files changed, 254 insertions(+), 233 deletions(-) diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index 841b197..272d17b 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -81,27 +81,6 @@ dependencies = [ "memchr", ] -[[package]] -name = "aide" -version = "0.15.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6966317188cdfe54c58c0900a195d021294afb3ece9b7073d09e4018dbb1e3a2" -dependencies = [ - "axum", - "bytes", - "cfg-if", - "http", - "indexmap", - "schemars 0.9.0", - "serde", - "serde_json", - "serde_qs 0.14.0", - "thiserror 2.0.18", - "tower-layer", - "tower-service", - "tracing", -] - [[package]] name = "aide" version = "0.16.0-alpha.4" @@ -113,10 +92,10 @@ dependencies = [ "cfg-if", "http", "indexmap", - "schemars 1.2.1", + "schemars", "serde", "serde_json", - "serde_qs 1.1.2", + "serde_qs", "thiserror 2.0.18", "tower-layer", "tower-service", @@ -311,31 +290,6 @@ dependencies = [ "tracing", ] -[[package]] -name = "axum-extra" -version = "0.12.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "be44683b41ccb9ab2d23a5230015c9c3c55be97a25e4428366de8873103f7970" -dependencies = [ - "axum", - "axum-core", - "bytes", - "form_urlencoded", - "futures-core", - "futures-util", - "http", - "http-body", - "http-body-util", - "mime", - "pin-project-lite", - "serde_core", - "serde_html_form", - "serde_path_to_error", - "tower-layer", - "tower-service", - "tracing", -] - [[package]] name = "axum-helmet" version = "1.0.2" @@ -349,17 +303,6 @@ dependencies = [ "tower-service", ] -[[package]] -name = "axum-macros" -version = "0.5.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7aa268c23bfbbd2c4363b9cd302a4f504fb2a9dfe7e3451d66f35dd392e20aca" -dependencies = [ - "proc-macro2", - "quote", - "syn", -] - [[package]] name = "axum-plugin" version = "0.2.0" @@ -382,30 +325,6 @@ dependencies = [ "syn", ] -[[package]] -name = "axum-typed-routing" -version = "0.4.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a61ccff37342bbcd7fc468ab6d96e9547a3c81bb2b027f108dda07e0757ab477" -dependencies = [ - "aide 0.15.1", - "axum", - "axum-extra", - "axum-macros", - "axum-typed-routing-macros", -] - -[[package]] -name = "axum-typed-routing-macros" -version = "0.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "299699055db19bfe910cb2275056062b6dbb0198f7d6d365e1f4f5e2f9f6138f" -dependencies = [ - "proc-macro2", - "quote", - "syn", -] - [[package]] name = "base64" version = "0.22.1" @@ -814,27 +733,6 @@ dependencies = [ "serde_core", ] -[[package]] -name = "derive_more" -version = "2.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d751e9e49156b02b44f9c1815bcb94b984cdcc4396ecc32521c739452808b134" -dependencies = [ - "derive_more-impl", -] - -[[package]] -name = "derive_more-impl" -version = "2.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "799a97264921d8623a957f6c3b9011f3b5492f557bbb7a5a19b7fa6d06ba8dcb" -dependencies = [ - "proc-macro2", - "quote", - "rustc_version", - "syn", -] - [[package]] name = "diesel" version = "2.3.10" @@ -2403,7 +2301,7 @@ name = "rs-chat-api" version = "0.1.0" dependencies = [ "aes-gcm 0.11.0", - "aide 0.16.0-alpha.4", + "aide", "aide-docs-macro", "anyhow", "async-stream", @@ -2411,9 +2309,7 @@ dependencies = [ "axum", "axum-helmet", "axum-plugin", - "axum-typed-routing", "chrono", - "derive_more", "diesel", "diesel-async", "diesel-derive-enum", @@ -2426,7 +2322,7 @@ dependencies = [ "hex", "reqwest", "reqwest-websocket", - "schemars 1.2.1", + "schemars", "serde", "serde_json", "serde_with", @@ -2567,20 +2463,6 @@ dependencies = [ "windows-sys 0.61.2", ] -[[package]] -name = "schemars" -version = "0.9.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4cd191f9397d57d581cddd31014772520aa448f65ef991055d7f61582c65165f" -dependencies = [ - "dyn-clone", - "indexmap", - "ref-cast", - "schemars_derive 0.9.0", - "serde", - "serde_json", -] - [[package]] name = "schemars" version = "1.2.1" @@ -2591,24 +2473,12 @@ dependencies = [ "dyn-clone", "indexmap", "ref-cast", - "schemars_derive 1.2.1", + "schemars_derive", "serde", "serde_json", "uuid", ] -[[package]] -name = "schemars_derive" -version = "0.9.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5016d94c77c6d32f0b8e08b781f7dc8a90c2007d4e77472cc2807bc10a8438fe" -dependencies = [ - "proc-macro2", - "quote", - "serde_derive_internals", - "syn", -] - [[package]] name = "schemars_derive" version = "1.2.1" @@ -2697,19 +2567,6 @@ dependencies = [ "syn", ] -[[package]] -name = "serde_html_form" -version = "0.2.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b2f2d7ff8a2140333718bb329f5c40fc5f0865b84c426183ce14c97d2ab8154f" -dependencies = [ - "form_urlencoded", - "indexmap", - "itoa", - "ryu", - "serde_core", -] - [[package]] name = "serde_json" version = "1.0.150" @@ -2734,19 +2591,6 @@ dependencies = [ "serde_core", ] -[[package]] -name = "serde_qs" -version = "0.14.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8b417bedc008acbdf6d6b4bc482d29859924114bbe2650b7921fb68a261d0aa6" -dependencies = [ - "axum", - "futures", - "percent-encoding", - "serde", - "thiserror 2.0.18", -] - [[package]] name = "serde_qs" version = "1.1.2" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 7738a41..82d7a48 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -22,13 +22,11 @@ axum-plugin = { git = "https://git.fasharp.io/fa-sharp/axum-plugin", rev = "9f72278b3c" } -axum-typed-routing = { version = "0.4.0", features = ["aide"] } chrono = { version = "0.4.45", default-features = false, features = ["now", "serde", "std"] } -derive_more = { version = "2.1.1", features = ["deref", "into"] } diesel = { version = "2.3.10", default-features = false, diff --git a/server-new/crates/aide-docs-macro/src/lib.rs b/server-new/crates/aide-docs-macro/src/lib.rs index 5c46194..54b77ea 100644 --- a/server-new/crates/aide-docs-macro/src/lib.rs +++ b/server-new/crates/aide-docs-macro/src/lib.rs @@ -1,6 +1,9 @@ use proc_macro::TokenStream; use quote::{format_ident, quote}; -use syn::{ItemFn, LitStr, Result, Token, parse::Parse, parse::ParseStream, parse_macro_input}; +use syn::{ + Expr, Ident, ItemFn, LitStr, Result, Token, Type, parse::Parse, parse::ParseStream, + parse_macro_input, +}; struct DocsArgs { summary: LitStr, @@ -62,3 +65,193 @@ pub fn docs(args: TokenStream, input: TokenStream) -> TokenStream { } .into() } + +struct ApiRoutes { + state: Type, + tag: Expr, + routes: Vec, +} + +struct ApiRoute { + methods: Vec, + path: LitStr, + handler: Ident, + summary: LitStr, + description: Option, +} + +struct RouteMethod { + ident: Ident, +} + +impl RouteMethod { + fn route_fn(&self) -> Result { + let method = self.ident.to_string(); + let fn_name = match method.as_str() { + "GET" => "get_with", + "POST" => "post_with", + "DELETE" => "delete_with", + _ => { + return Err(syn::Error::new_spanned( + &self.ident, + "expected one of GET, POST, DELETE", + )); + } + }; + + Ok(format_ident!("{fn_name}")) + } +} + +impl Parse for ApiRoutes { + fn parse(input: ParseStream<'_>) -> Result { + parse_label(input, "state")?; + input.parse::()?; + let state = input.parse()?; + input.parse::()?; + + parse_label(input, "tag")?; + input.parse::()?; + let tag = input.parse()?; + input.parse::()?; + + let mut routes = Vec::new(); + while !input.is_empty() { + routes.push(input.parse()?); + } + + Ok(Self { state, tag, routes }) + } +} + +impl Parse for ApiRoute { + fn parse(input: ParseStream<'_>) -> Result { + let mut methods = vec![RouteMethod { + ident: input.parse()?, + }]; + + while input.peek(Token![,]) { + let fork = input.fork(); + fork.parse::()?; + if fork.peek(Ident) { + input.parse::()?; + methods.push(RouteMethod { + ident: input.parse()?, + }); + } else { + break; + } + } + + let path = input.parse()?; + input.parse::]>()?; + let handler = input.parse()?; + input.parse::()?; + let summary = input.parse()?; + + let description = if input.peek(Token![,]) { + input.parse::()?; + Some(input.parse()?) + } else { + None + }; + + input.parse::()?; + + Ok(Self { + methods, + path, + handler, + summary, + description, + }) + } +} + +fn parse_label(input: ParseStream<'_>, expected: &str) -> Result<()> { + let label: Ident = input.parse()?; + if label == expected { + Ok(()) + } else { + Err(syn::Error::new_spanned( + label, + format!("expected `{expected}`"), + )) + } +} + +/// Generate a `routes()` function and the matching `_docs` functions. +/// +/// # Syntax +/// ```ignore +/// api_routes! { +/// state: AppState, +/// tag: ApiTag::Auth, +/// GET "/user" => get_user, "Get user", "Get the current user"; +/// GET, POST "/logout" => logout, "Logout"; +/// } +/// ``` +#[proc_macro] +pub fn api_routes(input: TokenStream) -> TokenStream { + let api_routes = parse_macro_input!(input as ApiRoutes); + let state = api_routes.state; + let tag = api_routes.tag; + + let docs_functions = api_routes.routes.iter().map(|route| { + let docs_name = format_ident!("{}_docs", route.handler); + let handler_name = route.handler.to_string(); + let summary = &route.summary; + let description = route.description.as_ref().map(|description| { + quote! { + .description(#description) + } + }); + + quote! { + fn #docs_name( + op: ::aide::transform::TransformOperation, + ) -> ::aide::transform::TransformOperation { + op.id(#handler_name) + .summary(#summary) + #description + } + } + }); + + let route_calls = api_routes.routes.iter().map(|route| { + let path = &route.path; + let handler = &route.handler; + let docs_name = format_ident!("{}_docs", route.handler); + + let mut methods = route + .methods + .iter() + .map(RouteMethod::route_fn) + .collect::>>()?; + let first_method = methods.remove(0); + + Ok(quote! { + .api_route( + #path, + ::aide::axum::routing::#first_method(#handler, #docs_name) + #(.#methods(#handler, #docs_name))* + ) + }) + }); + + let route_calls = match route_calls.collect::>>() { + Ok(route_calls) => route_calls, + Err(error) => return error.into_compile_error().into(), + }; + + quote! { + #(#docs_functions)* + + pub fn routes() -> ::aide::axum::ApiRouter<#state> { + ::aide::axum::ApiRouter::new() + #(#route_calls)* + .with_path_items(|op| op.tag(#tag.into())) + } + } + .into() +} diff --git a/server-new/src/api/api_key.rs b/server-new/src/api/api_key.rs index 252b4c8..a116932 100644 --- a/server-new/src/api/api_key.rs +++ b/server-new/src/api/api_key.rs @@ -1,8 +1,4 @@ -use aide::axum::{ - ApiRouter, - routing::{delete_with, get_with, post_with}, -}; -use aide_docs_macro::docs; +use aide_docs_macro::api_routes; use axum::{ Json, extract::{Path, State}, @@ -20,15 +16,14 @@ use crate::{ state::AppState, }; -pub fn routes() -> ApiRouter { - ApiRouter::new() - .api_route("/", get_with(list_api_keys, list_api_keys_docs)) - .api_route("/", post_with(create_api_key, create_api_key_docs)) - .api_route("/{id}", delete_with(delete_api_key, delete_api_key_docs)) - .with_path_items(|op| op.tag(ApiTag::ApiKey.into())) +api_routes! { + state: AppState, + tag: ApiTag::ApiKey, + GET "/" => list_api_keys, "List API keys"; + POST "/" => create_api_key, "Create API key"; + DELETE "/{id}" => delete_api_key, "Delete API key"; } -#[docs("List API keys")] async fn list_api_keys( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, @@ -42,7 +37,6 @@ struct ApiKeyCreateInput { name: String, } -#[docs("Create API key")] async fn create_api_key( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, @@ -63,7 +57,6 @@ struct ApiKeyCreateResponse { key: String, } -#[docs("Delete API key")] async fn delete_api_key( Path(id): Path, CurrentUser { user_id }: CurrentUser, diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index 23f836e..6c2b990 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -1,5 +1,4 @@ -use aide::axum::{ApiRouter, routing::get_with}; -use aide_docs_macro::docs; +use aide_docs_macro::api_routes; use axum::{ Extension, Json, extract::{Path, Query, State}, @@ -18,23 +17,16 @@ use crate::{ state::AppState, }; -pub fn routes() -> ApiRouter { - ApiRouter::new() - .api_route("/user", get_with(get_user, get_user_docs)) - .api_route("/config", get_with(get_config, get_config_docs)) - .api_route("/login", get_with(oauth_login, oauth_login_docs)) - .api_route( - "/login/callback", - get_with(oauth_callback, oauth_callback_docs), - ) - .api_route( - "/logout", - get_with(logout, logout_docs).post_with(logout, logout_docs), - ) - .with_path_items(|op| op.tag(ApiTag::Auth.into())) +api_routes! { + state: AppState, + tag: ApiTag::Auth, + GET "/user" => get_user, "Get user", "Get the current user"; + GET "/config" => get_config, "Get auth config", "Get the current auth configuration of the server"; + GET "/login/{provider}" => oauth_login, "OAuth login", "OAuth login redirect"; + GET "/login/{provider}/callback" => oauth_callback, "OAuth login callback"; + GET, POST "/logout" => logout, "Logout"; } -#[docs("Get user", "Get the current user")] async fn get_user( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, @@ -44,7 +36,6 @@ async fn get_user( Ok(Json(user)) } -#[docs("Get auth config", "Get the current auth configuration of the server")] async fn get_config(auth_config: PublicAuthConfig) -> Json { Json(auth_config) } @@ -53,7 +44,6 @@ fn oauth_callback_path(route_prefix: &'static str, provider: &OAuthProviderEnum) format!("{route_prefix}/login/{provider}/callback") } -#[docs("OAuth login", "OAuth login redirect")] async fn oauth_login( Path(provider): Path, Extension(RoutePrefix(prefix)): Extension, @@ -74,7 +64,6 @@ struct OAuthCallbackQuery { state: String, } -#[docs("OAuth login callback sdf")] async fn oauth_callback( Path(provider): Path, Query(query): Query, @@ -106,7 +95,6 @@ async fn oauth_callback( Ok(Redirect::to(&app_state.config.server.base_url)) } -#[docs("Logout")] async fn logout( AppSession { session, .. }: AppSession, State(state): State, diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index eabcc8c..721f93b 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -1,5 +1,4 @@ -use aide::axum::{ApiRouter, routing::post_with}; -use aide_docs_macro::docs; +use aide_docs_macro::api_routes; use axum::{ Json, extract::{Path, State}, @@ -9,19 +8,20 @@ use serde::{Deserialize, Serialize}; use uuid::Uuid; use crate::{ + api::ApiTag, error::{AppError, AppResult}, extractors::{CurrentUser, Database}, llm::types::{LlmChatOptions, LlmUserMessage}, state::AppState, }; -pub fn routes() -> ApiRouter { - ApiRouter::new() - .api_route("/prompt", post_with(prompt, prompt_docs)) - .api_route( - "/session/{session_id}", - post_with(chat_stream, chat_stream_docs), - ) +api_routes! { + state: AppState, + tag: ApiTag::Chat, + POST "/prompt" => prompt, "Prompt", + "Send a single prompt to a provider and get the response"; + POST "/session/{session_id}" => chat_stream, "Chat", + "Send a message in a chat session and stream the response"; } #[derive(Debug, Deserialize, JsonSchema)] @@ -34,7 +34,6 @@ struct PromptInput { options: LlmChatOptions, } -#[docs("Prompt", "Send a single prompt to a provider and get the response")] async fn prompt( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, @@ -75,10 +74,6 @@ struct ChatInput { options: LlmChatOptions, } -#[docs( - "Send chat", - "Send a message in a chat session and stream the response" -)] async fn chat_stream( Path(session_id): Path, CurrentUser { user_id }: CurrentUser, diff --git a/server-new/src/api/health.rs b/server-new/src/api/health.rs index 65671ea..61f08d6 100644 --- a/server-new/src/api/health.rs +++ b/server-new/src/api/health.rs @@ -1,13 +1,13 @@ -use aide::axum::ApiRouter; -use axum_typed_routing::{TypedApiRouter, api_route}; +use aide::axum::{ApiRouter, routing::get_with}; +use aide_docs_macro::docs; use crate::state::AppState; pub fn routes() -> ApiRouter { - ApiRouter::new().typed_api_route(health) + ApiRouter::new().api_route("/", get_with(health, health_docs)) } -#[api_route(GET "/" with AppState { summary: "Health route" })] +#[docs("Health route")] async fn health() -> &'static str { "OK" } diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index 3803fc3..4a64267 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -1,7 +1,11 @@ use std::sync::Arc; -use aide::{axum::ApiRouter, openapi::OpenApi, swagger::Swagger}; -use axum::{Extension, http::header, response::IntoResponse, routing::get}; +use aide::{ + axum::ApiRouter, + openapi::{OpenApi, Server}, + swagger::Swagger, +}; +use axum::{Extension, routing::get}; use axum_plugin::AdHocPlugin; use strum::{AsRefStr, Display, EnumIter, EnumMessage, IntoEnumIterator, IntoStaticStr}; @@ -26,14 +30,6 @@ enum ApiTag { pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("API routes").on_setup(|router, _state| { let mut openapi = OpenApi::default(); - for tag in ApiTag::iter() { - openapi.tags.push(aide::openapi::Tag { - name: tag.to_string(), - description: tag.get_message().map(String::from), - ..Default::default() - }) - } - let api_routes = ApiRouter::new() .nest("/api_key", api_key::routes()) .nest( @@ -42,12 +38,30 @@ pub fn plugin() -> AdHocPlugin { ) .nest("/chat", chat::routes()) .nest("/health", health::routes()) - .finish_api(&mut openapi); + .finish_api_with(&mut openapi, |op| { + let mut op = op + .title("RsChat API") + .description("OpenAPI specification for the RsChat server") + .server(Server { + url: String::from("/api"), + ..Default::default() + }); + for tag in ApiTag::iter() { + op = op.tag(aide::openapi::Tag { + name: tag.to_string(), + description: tag.get_message().map(String::from), + ..Default::default() + }); + } + + op + }); let api_routes_with_docs = api_routes .route( "/docs/openapi.json", - get(openapi_route).layer(Extension(Arc::new(openapi))), + get(async |Extension(openapi): Extension>| axum::Json(openapi)) + .layer(Extension(Arc::new(openapi))), ) .route( "/docs", @@ -60,10 +74,6 @@ pub fn plugin() -> AdHocPlugin { }) } -async fn openapi_route(Extension(openapi): Extension>) -> impl IntoResponse { - axum::Json(openapi) -} - /// Extension to pass the route prefix to child routes #[derive(Clone)] struct RoutePrefix(&'static str); From 433fabbbf4cad2dbee54d244604cb2e25e61b3e0 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sun, 5 Jul 2026 23:29:17 -0400 Subject: [PATCH 064/111] api doc tweaks --- server-new/crates/aide-docs-macro/src/lib.rs | 96 +++++++++++++++----- server-new/src/api/auth.rs | 4 +- 2 files changed, 77 insertions(+), 23 deletions(-) diff --git a/server-new/crates/aide-docs-macro/src/lib.rs b/server-new/crates/aide-docs-macro/src/lib.rs index 54b77ea..d23b72b 100644 --- a/server-new/crates/aide-docs-macro/src/lib.rs +++ b/server-new/crates/aide-docs-macro/src/lib.rs @@ -68,7 +68,7 @@ pub fn docs(args: TokenStream, input: TokenStream) -> TokenStream { struct ApiRoutes { state: Type, - tag: Expr, + tag: Option, routes: Vec, } @@ -85,9 +85,16 @@ struct RouteMethod { } impl RouteMethod { + fn name(&self) -> String { + self.ident.to_string() + } + + fn operation_prefix(&self) -> String { + self.name().to_lowercase() + } + fn route_fn(&self) -> Result { - let method = self.ident.to_string(); - let fn_name = match method.as_str() { + let fn_name = match self.name().as_str() { "GET" => "get_with", "POST" => "post_with", "DELETE" => "delete_with", @@ -110,10 +117,15 @@ impl Parse for ApiRoutes { let state = input.parse()?; input.parse::()?; - parse_label(input, "tag")?; - input.parse::()?; - let tag = input.parse()?; - input.parse::()?; + let tag = if next_label_is(input, "tag") { + parse_label(input, "tag")?; + input.parse::()?; + let tag = input.parse()?; + input.parse::()?; + Some(tag) + } else { + None + }; let mut routes = Vec::new(); while !input.is_empty() { @@ -180,13 +192,18 @@ fn parse_label(input: ParseStream<'_>, expected: &str) -> Result<()> { } } +fn next_label_is(input: ParseStream<'_>, expected: &str) -> bool { + let fork = input.fork(); + fork.parse::().is_ok_and(|label| label == expected) && fork.peek(Token![:]) +} + /// Generate a `routes()` function and the matching `_docs` functions. /// /// # Syntax /// ```ignore /// api_routes! { /// state: AppState, -/// tag: ApiTag::Auth, +/// tag: ApiTag::Auth, // optional /// GET "/user" => get_user, "Get user", "Get the current user"; /// GET, POST "/logout" => logout, "Logout"; /// } @@ -195,46 +212,60 @@ fn parse_label(input: ParseStream<'_>, expected: &str) -> Result<()> { pub fn api_routes(input: TokenStream) -> TokenStream { let api_routes = parse_macro_input!(input as ApiRoutes); let state = api_routes.state; - let tag = api_routes.tag; - let docs_functions = api_routes.routes.iter().map(|route| { - let docs_name = format_ident!("{}_docs", route.handler); - let handler_name = route.handler.to_string(); + let docs_functions = api_routes.routes.iter().flat_map(|route| { let summary = &route.summary; let description = route.description.as_ref().map(|description| { quote! { .description(#description) } }); + let multiple_methods = route.methods.len() > 1; - quote! { + route.methods.iter().map(move |method| { + let docs_name = docs_name(route, method, multiple_methods); + let operation_id = operation_id(route, method, multiple_methods); + + quote! { fn #docs_name( op: ::aide::transform::TransformOperation, ) -> ::aide::transform::TransformOperation { - op.id(#handler_name) + op.id(#operation_id) .summary(#summary) #description } - } + } + }) }); let route_calls = api_routes.routes.iter().map(|route| { let path = &route.path; let handler = &route.handler; - let docs_name = format_ident!("{}_docs", route.handler); + let multiple_methods = route.methods.len() > 1; let mut methods = route .methods .iter() - .map(RouteMethod::route_fn) + .map(|method| { + let route_fn = method.route_fn()?; + let docs_name = docs_name(route, method, multiple_methods); + + Ok((route_fn, docs_name)) + }) .collect::>>()?; - let first_method = methods.remove(0); + let (first_method, first_docs_name) = methods.remove(0); + + let additional_methods = methods.into_iter().map(|(method, docs_name)| { + quote! { + .#method(#handler, #docs_name) + } + }); Ok(quote! { .api_route( #path, - ::aide::axum::routing::#first_method(#handler, #docs_name) - #(.#methods(#handler, #docs_name))* + ::aide::axum::routing::#first_method(#handler, #first_docs_name) + #(#additional_methods)* ) }) }); @@ -244,14 +275,37 @@ pub fn api_routes(input: TokenStream) -> TokenStream { Err(error) => return error.into_compile_error().into(), }; + let tag = api_routes.tag.map(|tag| { + quote! { + .with_path_items(|op| op.tag(#tag.into())) + } + }); + quote! { #(#docs_functions)* pub fn routes() -> ::aide::axum::ApiRouter<#state> { ::aide::axum::ApiRouter::new() #(#route_calls)* - .with_path_items(|op| op.tag(#tag.into())) + #tag } } .into() } + +fn docs_name(route: &ApiRoute, method: &RouteMethod, multiple_methods: bool) -> Ident { + if multiple_methods { + let method = method.operation_prefix(); + format_ident!("{}_{}_docs", method, route.handler) + } else { + format_ident!("{}_docs", route.handler) + } +} + +fn operation_id(route: &ApiRoute, method: &RouteMethod, multiple_methods: bool) -> String { + if multiple_methods { + format!("{}_{}", method.operation_prefix(), route.handler) + } else { + route.handler.to_string() + } +} diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index 6c2b990..f3939c6 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -21,7 +21,7 @@ api_routes! { state: AppState, tag: ApiTag::Auth, GET "/user" => get_user, "Get user", "Get the current user"; - GET "/config" => get_config, "Get auth config", "Get the current auth configuration of the server"; + GET "/config" => get_auth_config, "Get auth config", "Get the current auth configuration of the server"; GET "/login/{provider}" => oauth_login, "OAuth login", "OAuth login redirect"; GET "/login/{provider}/callback" => oauth_callback, "OAuth login callback"; GET, POST "/logout" => logout, "Logout"; @@ -36,7 +36,7 @@ async fn get_user( Ok(Json(user)) } -async fn get_config(auth_config: PublicAuthConfig) -> Json { +async fn get_auth_config(auth_config: PublicAuthConfig) -> Json { Json(auth_config) } From 8b08455fa7abe59bbeabf3071efdc43ad0b7024c Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sun, 5 Jul 2026 23:57:17 -0400 Subject: [PATCH 065/111] api doc: add errors --- server-new/crates/aide-docs-macro/Cargo.toml | 2 +- server-new/src/error.rs | 37 ++++++++++++++++---- 2 files changed, 32 insertions(+), 7 deletions(-) diff --git a/server-new/crates/aide-docs-macro/Cargo.toml b/server-new/crates/aide-docs-macro/Cargo.toml index b655ebc..d01c3ff 100644 --- a/server-new/crates/aide-docs-macro/Cargo.toml +++ b/server-new/crates/aide-docs-macro/Cargo.toml @@ -2,7 +2,7 @@ name = "aide-docs-macro" version = "0.1.0" edition = "2024" -description = "Internal attribute macro for API operation docs" +description = "Convenience macros for aide API docs" publish = false [lib] diff --git a/server-new/src/error.rs b/server-new/src/error.rs index ef4e324..57eda5f 100644 --- a/server-new/src/error.rs +++ b/server-new/src/error.rs @@ -1,9 +1,10 @@ -use aide::OperationIo; +use aide::OperationOutput; use axum::{ Json, http::StatusCode, response::{IntoResponse, Response}, }; +use schemars::JsonSchema; use serde::Serialize; use crate::db::DbPoolError; @@ -12,7 +13,7 @@ use crate::db::DbPoolError; pub type AppResult = Result; /// Global API error type -#[derive(Debug, OperationIo)] +#[derive(Debug)] pub struct AppError { status: StatusCode, message: String, @@ -64,13 +65,13 @@ impl From for AppError { } } -#[derive(Debug, Serialize)] -struct ErrorResponse { +#[derive(Debug, Serialize, JsonSchema)] +pub struct ErrorResponse { error: ErrorBody, } -#[derive(Debug, Serialize)] -struct ErrorBody { +#[derive(Debug, Serialize, JsonSchema)] +pub struct ErrorBody { message: String, status: u16, } @@ -91,3 +92,27 @@ impl IntoResponse for AppError { (self.status, Json(response)).into_response() } } + +impl OperationOutput for AppError { + type Inner = ErrorResponse; + + fn inferred_responses( + ctx: &mut aide::generate::GenContext, + operation: &mut aide::openapi::Operation, + ) -> Vec<(Option, aide::openapi::Response)> { + if let Some(response) = Json::::operation_response(ctx, operation) { + let status_codes = [ + StatusCode::BAD_REQUEST, + StatusCode::UNAUTHORIZED, + StatusCode::NOT_FOUND, + StatusCode::INTERNAL_SERVER_ERROR, + ]; + Vec::from_iter(status_codes.into_iter().map(|code| { + let aide_code = aide::openapi::StatusCode::Code(code.as_u16()); + (Some(aide_code), response.clone()) + })) + } else { + Vec::new() + } + } +} From cd3b71d18a0eb14512248768a4f8e3aff8b62ee7 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 6 Jul 2026 00:46:48 -0400 Subject: [PATCH 066/111] api doc: add api key auth --- server-new/Cargo.lock | 1 + server-new/Cargo.toml | 5 ++++- server-new/crates/aide-docs-macro/src/lib.rs | 10 +++++----- server-new/src/api/health.rs | 4 ++-- server-new/src/api/mod.rs | 13 +++++++++++-- server-new/src/extractors/user.rs | 16 ++++++++++++++-- 6 files changed, 37 insertions(+), 12 deletions(-) diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index 272d17b..fe8560f 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -2573,6 +2573,7 @@ version = "1.0.150" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e8014e44b4736ed0538adeecded0fce2a272f22dc9578a7eb6b2d9993c74cfb9" dependencies = [ + "indexmap", "itoa", "memchr", "serde", diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 82d7a48..a60526f 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -54,7 +54,10 @@ reqwest = { features = ["default-tls", "json", "stream"] } reqwest-websocket = { version = "0.6.0", features = ["json"] } -schemars = { version = "1.2.1", features = ["chrono04", "uuid1"] } +schemars = { + version = "1.2.1", + features = ["chrono04", "preserve_order", "uuid1"] +} serde = { version = "1.0.228", features = ["derive"] } serde_json = "1.0.150" serde_with = { diff --git a/server-new/crates/aide-docs-macro/src/lib.rs b/server-new/crates/aide-docs-macro/src/lib.rs index d23b72b..d043b7f 100644 --- a/server-new/crates/aide-docs-macro/src/lib.rs +++ b/server-new/crates/aide-docs-macro/src/lib.rs @@ -31,14 +31,14 @@ impl Parse for DocsArgs { } } -/// Convenience macro for generating API docs with aide. Generates a function +/// Convenience macro for generating API docs for the route handler. Generates a function /// called `_docs` that can be passed as the transform function -/// to `get_with`, `post_with`, etc. +/// to aide's `get_with`, `post_with`, etc. /// /// # Syntax -/// `#[docs("", ""]` +/// `#[handler_docs("" (, "")]` #[proc_macro_attribute] -pub fn docs(args: TokenStream, input: TokenStream) -> TokenStream { +pub fn handler_docs(args: TokenStream, input: TokenStream) -> TokenStream { let args = parse_macro_input!(args as DocsArgs); let handler = parse_macro_input!(input as ItemFn); let handler_name = &handler.sig.ident; @@ -197,7 +197,7 @@ fn next_label_is(input: ParseStream<'_>, expected: &str) -> bool { fork.parse::().is_ok_and(|label| label == expected) && fork.peek(Token![:]) } -/// Generate a `routes()` function and the matching `_docs` functions. +/// Generate a `routes()` function with attached API docs /// /// # Syntax /// ```ignore diff --git a/server-new/src/api/health.rs b/server-new/src/api/health.rs index 61f08d6..669c571 100644 --- a/server-new/src/api/health.rs +++ b/server-new/src/api/health.rs @@ -1,5 +1,5 @@ use aide::axum::{ApiRouter, routing::get_with}; -use aide_docs_macro::docs; +use aide_docs_macro::handler_docs; use crate::state::AppState; @@ -7,7 +7,7 @@ pub fn routes() -> ApiRouter { ApiRouter::new().api_route("/", get_with(health, health_docs)) } -#[docs("Health route")] +#[handler_docs("Health route")] async fn health() -> &'static str { "OK" } diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index 4a64267..f12b238 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -2,7 +2,7 @@ use std::sync::Arc; use aide::{ axum::ApiRouter, - openapi::{OpenApi, Server}, + openapi::{OpenApi, SecurityScheme, Server}, swagger::Swagger, }; use axum::{Extension, routing::get}; @@ -45,7 +45,16 @@ pub fn plugin() -> AdHocPlugin { .server(Server { url: String::from("/api"), ..Default::default() - }); + }) + .security_scheme( + "ApiKey", + SecurityScheme::Http { + scheme: String::from("bearer"), + bearer_format: Some(String::from("bearer")), + description: Some(String::from("RsChat API key")), + extensions: Default::default(), + }, + ); for tag in ApiTag::iter() { op = op.tag(aide::openapi::Tag { name: tag.to_string(), diff --git a/server-new/src/extractors/user.rs b/server-new/src/extractors/user.rs index ea9e473..5955120 100644 --- a/server-new/src/extractors/user.rs +++ b/server-new/src/extractors/user.rs @@ -1,4 +1,4 @@ -use aide::OperationIo; +use aide::OperationInput; use anyhow::anyhow; use axum::{ extract::{FromRequestParts, OptionalFromRequestParts}, @@ -17,7 +17,7 @@ if there is no active user. - If used as `Option`, will be `Some` if there is an active user and `None` otherwise. */ -#[derive(Debug, Clone, Serialize, Deserialize, OperationIo)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct CurrentUser { pub user_id: Uuid, } @@ -93,3 +93,15 @@ impl FromRequestParts for CurrentUser { } } } + +impl OperationInput for CurrentUser { + fn operation_input( + _ctx: &mut aide::generate::GenContext, + operation: &mut aide::openapi::Operation, + ) { + let security_reqs = [(String::from("ApiKey"), vec![])]; + operation + .security + .push(FromIterator::from_iter(security_reqs)) + } +} From 6aa5d45f87a4e8a9ee4c61f7014baebcc9030ac0 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 6 Jul 2026 00:59:50 -0400 Subject: [PATCH 067/111] api doc: tweaks --- server-new/src/api/mod.rs | 75 ++++++++++++++++++------------- server-new/src/extractors/user.rs | 2 +- 2 files changed, 44 insertions(+), 33 deletions(-) diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index f12b238..e99e8b2 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -16,6 +16,10 @@ pub mod auth; pub mod chat; pub mod health; +const API_BASE: &str = "/api/v1"; +const API_AUTH_BASE: &str = "/api/v1/auth"; +pub const API_KEY_SCHEME: &str = "ApiKey"; + #[derive(Display, AsRefStr, IntoStaticStr, EnumMessage, EnumIter)] enum ApiTag { #[strum(message = "Manage API keys")] @@ -24,9 +28,11 @@ enum ApiTag { Auth, #[strum(message = "Chats and sessions")] Chat, + #[strum(message = "AI / LLM Providers")] + Provider, } -/// Adds all API routes with OpenAPI docs to the server under `/api` +/// Adds all API routes with OpenAPI docs to the server under `/api/v1` pub fn plugin() -> AdHocPlugin { AdHocPlugin::named("API routes").on_setup(|router, _state| { let mut openapi = OpenApi::default(); @@ -34,37 +40,11 @@ pub fn plugin() -> AdHocPlugin { .nest("/api_key", api_key::routes()) .nest( "/auth", - auth::routes().layer(Extension(RoutePrefix("/api/auth"))), + auth::routes().layer(Extension(RoutePrefix(API_AUTH_BASE))), ) .nest("/chat", chat::routes()) .nest("/health", health::routes()) - .finish_api_with(&mut openapi, |op| { - let mut op = op - .title("RsChat API") - .description("OpenAPI specification for the RsChat server") - .server(Server { - url: String::from("/api"), - ..Default::default() - }) - .security_scheme( - "ApiKey", - SecurityScheme::Http { - scheme: String::from("bearer"), - bearer_format: Some(String::from("bearer")), - description: Some(String::from("RsChat API key")), - extensions: Default::default(), - }, - ); - for tag in ApiTag::iter() { - op = op.tag(aide::openapi::Tag { - name: tag.to_string(), - description: tag.get_message().map(String::from), - ..Default::default() - }); - } - - op - }); + .finish_api_with(&mut openapi, build_openapi_doc); let api_routes_with_docs = api_routes .route( @@ -74,15 +54,46 @@ pub fn plugin() -> AdHocPlugin { ) .route( "/docs", - get(Swagger::new("/api/docs/openapi.json") - .with_title("RsChat API") + get(Swagger::new(format!("{API_BASE}/docs/openapi.json")) + .with_title("RsChat API documentation") .axum_handler()), ); - Ok(router.nest("/api", api_routes_with_docs)) + Ok(router.nest(API_BASE, api_routes_with_docs)) }) } /// Extension to pass the route prefix to child routes #[derive(Clone)] struct RoutePrefix(&'static str); + +/// Build the OpenAPI docs +fn build_openapi_doc( + op: aide::transform::TransformOpenApi<'_>, +) -> aide::transform::TransformOpenApi<'_> { + let mut op = op + .title("RsChat API") + .description("OpenAPI specification for the RsChat server") + .server(Server { + url: String::from(API_BASE), + ..Default::default() + }) + .security_scheme( + API_KEY_SCHEME, + SecurityScheme::Http { + scheme: String::from("bearer"), + bearer_format: Some(String::from("bearer")), + description: Some(String::from("RsChat API key")), + extensions: Default::default(), + }, + ); + for tag in ApiTag::iter() { + op = op.tag(aide::openapi::Tag { + name: tag.to_string(), + description: tag.get_message().map(String::from), + ..Default::default() + }); + } + + op +} diff --git a/server-new/src/extractors/user.rs b/server-new/src/extractors/user.rs index 5955120..2b47ef1 100644 --- a/server-new/src/extractors/user.rs +++ b/server-new/src/extractors/user.rs @@ -99,7 +99,7 @@ impl OperationInput for CurrentUser { _ctx: &mut aide::generate::GenContext, operation: &mut aide::openapi::Operation, ) { - let security_reqs = [(String::from("ApiKey"), vec![])]; + let security_reqs = [(String::from(crate::api::API_KEY_SCHEME), vec![])]; operation .security .push(FromIterator::from_iter(security_reqs)) From 2333e9c938c83e1d2bd349a695c0f3722cb2ea54 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 6 Jul 2026 01:14:36 -0400 Subject: [PATCH 068/111] api doc: add responses to macro --- server-new/crates/aide-docs-macro/src/lib.rs | 75 +++++++++++++++++++- server-new/src/api/auth.rs | 12 +++- 2 files changed, 81 insertions(+), 6 deletions(-) diff --git a/server-new/crates/aide-docs-macro/src/lib.rs b/server-new/crates/aide-docs-macro/src/lib.rs index d043b7f..bb2530f 100644 --- a/server-new/crates/aide-docs-macro/src/lib.rs +++ b/server-new/crates/aide-docs-macro/src/lib.rs @@ -1,8 +1,8 @@ use proc_macro::TokenStream; use quote::{format_ident, quote}; use syn::{ - Expr, Ident, ItemFn, LitStr, Result, Token, Type, parse::Parse, parse::ParseStream, - parse_macro_input, + Expr, Ident, ItemFn, LitInt, LitStr, Result, Token, Type, braced, parse::Parse, + parse::ParseStream, parse_macro_input, }; struct DocsArgs { @@ -78,6 +78,12 @@ struct ApiRoute { handler: Ident, summary: LitStr, description: Option, + responses: Vec, +} + +struct RouteResponse { + status: LitInt, + ty: Type, } struct RouteMethod { @@ -168,6 +174,12 @@ impl Parse for ApiRoute { None }; + let responses = if input.peek(syn::token::Brace) { + input.parse::()?.responses + } else { + Vec::new() + }; + input.parse::()?; Ok(Self { @@ -176,10 +188,54 @@ impl Parse for ApiRoute { handler, summary, description, + responses, + }) + } +} + +struct RouteOptions { + responses: Vec, +} + +impl Parse for RouteOptions { + fn parse(input: ParseStream<'_>) -> Result { + let content; + braced!(content in input); + + parse_label(&content, "responses")?; + content.parse::()?; + + let responses; + braced!(responses in content); + + let mut route_responses = Vec::new(); + while !responses.is_empty() { + route_responses.push(responses.parse()?); + if responses.peek(Token![,]) { + responses.parse::()?; + } + } + + if !content.is_empty() { + content.parse::()?; + } + + Ok(Self { + responses: route_responses, }) } } +impl Parse for RouteResponse { + fn parse(input: ParseStream<'_>) -> Result { + let status = input.parse()?; + input.parse::()?; + let ty = input.parse()?; + + Ok(Self { status, ty }) + } +} + fn parse_label(input: ParseStream<'_>, expected: &str) -> Result<()> { let label: Ident = input.parse()?; if label == expected { @@ -205,7 +261,7 @@ fn next_label_is(input: ParseStream<'_>, expected: &str) -> bool { /// state: AppState, /// tag: ApiTag::Auth, // optional /// GET "/user" => get_user, "Get user", "Get the current user"; -/// GET, POST "/logout" => logout, "Logout"; +/// GET, POST "/logout" => logout, "Logout" { responses: { 204: () } }; /// } /// ``` #[proc_macro] @@ -220,6 +276,18 @@ pub fn api_routes(input: TokenStream) -> TokenStream { .description(#description) } }); + let responses = route + .responses + .iter() + .map(|response| { + let status = &response.status; + let ty = &response.ty; + + quote! { + .response::<#status, #ty>() + } + }) + .collect::>(); let multiple_methods = route.methods.len() > 1; route.methods.iter().map(move |method| { @@ -233,6 +301,7 @@ pub fn api_routes(input: TokenStream) -> TokenStream { op.id(#operation_id) .summary(#summary) #description + #(#responses)* } } }) diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index f3939c6..6a61cdb 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -22,9 +22,15 @@ api_routes! { tag: ApiTag::Auth, GET "/user" => get_user, "Get user", "Get the current user"; GET "/config" => get_auth_config, "Get auth config", "Get the current auth configuration of the server"; - GET "/login/{provider}" => oauth_login, "OAuth login", "OAuth login redirect"; - GET "/login/{provider}/callback" => oauth_callback, "OAuth login callback"; - GET, POST "/logout" => logout, "Logout"; + GET "/login/{provider}" => oauth_login, "OAuth login", "OAuth login redirect" { + responses: { 303: () } + }; + GET "/login/{provider}/callback" => oauth_callback, "OAuth login callback" { + responses: { 303: () } + }; + GET, POST "/logout" => logout, "Logout" { + responses: { 204: () } + }; } async fn get_user( From 9176fe46075d6504551e40845bcb53cea00f8b19 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 6 Jul 2026 01:26:37 -0400 Subject: [PATCH 069/111] api doc: more macro tweaks --- server-new/crates/aide-docs-macro/src/lib.rs | 115 ++++++++++++++----- server-new/src/api/auth.rs | 12 +- server-new/src/api/chat.rs | 10 +- 3 files changed, 98 insertions(+), 39 deletions(-) diff --git a/server-new/crates/aide-docs-macro/src/lib.rs b/server-new/crates/aide-docs-macro/src/lib.rs index bb2530f..f0bd4e4 100644 --- a/server-new/crates/aide-docs-macro/src/lib.rs +++ b/server-new/crates/aide-docs-macro/src/lib.rs @@ -76,7 +76,7 @@ struct ApiRoute { methods: Vec, path: LitStr, handler: Ident, - summary: LitStr, + summary: Option, description: Option, responses: Vec, } @@ -165,64 +165,107 @@ impl Parse for ApiRoute { input.parse::]>()?; let handler = input.parse()?; input.parse::()?; - let summary = input.parse()?; - let description = if input.peek(Token![,]) { - input.parse::()?; + let summary = if input.peek(LitStr) { Some(input.parse()?) } else { None }; - let responses = if input.peek(syn::token::Brace) { - input.parse::()?.responses + let options = if input.peek(Token![,]) { + input.parse::()?; + if !input.peek(syn::token::Brace) { + return Err(input.error("expected route options block after summary comma")); + } + input.parse()? + } else if input.peek(syn::token::Brace) { + input.parse()? } else { - Vec::new() + RouteOptions::default() }; input.parse::()?; + if summary.is_none() && options.is_empty() { + return Err(input.error("expected summary string or route options block")); + } + Ok(Self { methods, path, handler, summary, - description, - responses, + description: options.description, + responses: options.responses, }) } } +#[derive(Default)] struct RouteOptions { + description: Option, responses: Vec, } +impl RouteOptions { + fn is_empty(&self) -> bool { + self.description.is_none() && self.responses.is_empty() + } +} + impl Parse for RouteOptions { fn parse(input: ParseStream<'_>) -> Result { let content; braced!(content in input); - parse_label(&content, "responses")?; - content.parse::()?; - - let responses; - braced!(responses in content); - - let mut route_responses = Vec::new(); - while !responses.is_empty() { - route_responses.push(responses.parse()?); - if responses.peek(Token![,]) { - responses.parse::()?; + let mut options = RouteOptions::default(); + while !content.is_empty() { + let label: Ident = content.parse()?; + content.parse::()?; + + match label.to_string().as_str() { + "description" => { + if options.description.is_some() { + return Err(syn::Error::new_spanned( + label, + "`description` can only be provided once", + )); + } + + options.description = Some(content.parse()?); + } + "responses" => { + if !options.responses.is_empty() { + return Err(syn::Error::new_spanned( + label, + "`responses` can only be provided once", + )); + } + + let responses; + braced!(responses in content); + + while !responses.is_empty() { + options.responses.push(responses.parse()?); + if responses.peek(Token![,]) { + responses.parse::()?; + } + } + } + _ => { + return Err(syn::Error::new_spanned( + label, + "expected `description` or `responses`", + )); + } } - } - if !content.is_empty() { - content.parse::()?; + if content.peek(Token![,]) { + content.parse::()?; + } } - Ok(Self { - responses: route_responses, - }) + Ok(options) } } @@ -260,8 +303,16 @@ fn next_label_is(input: ParseStream<'_>, expected: &str) -> bool { /// api_routes! { /// state: AppState, /// tag: ApiTag::Auth, // optional -/// GET "/user" => get_user, "Get user", "Get the current user"; -/// GET, POST "/logout" => logout, "Logout" { responses: { 204: () } }; +/// GET "/user" => get_user, "Get user", { +/// description: "Get the current user" +/// }; +/// GET, POST "/logout" => logout, "Logout", { +/// responses: { 204: () } +/// }; +/// POST "/sessions" => create_session, { +/// description: "Create a new session", +/// responses: { 201: Session } +/// }; /// } /// ``` #[proc_macro] @@ -270,7 +321,11 @@ pub fn api_routes(input: TokenStream) -> TokenStream { let state = api_routes.state; let docs_functions = api_routes.routes.iter().flat_map(|route| { - let summary = &route.summary; + let summary = route.summary.as_ref().map(|summary| { + quote! { + .summary(#summary) + } + }); let description = route.description.as_ref().map(|description| { quote! { .description(#description) @@ -299,7 +354,7 @@ pub fn api_routes(input: TokenStream) -> TokenStream { op: ::aide::transform::TransformOperation, ) -> ::aide::transform::TransformOperation { op.id(#operation_id) - .summary(#summary) + #summary #description #(#responses)* } diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index 6a61cdb..df88bee 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -20,15 +20,17 @@ use crate::{ api_routes! { state: AppState, tag: ApiTag::Auth, - GET "/user" => get_user, "Get user", "Get the current user"; - GET "/config" => get_auth_config, "Get auth config", "Get the current auth configuration of the server"; - GET "/login/{provider}" => oauth_login, "OAuth login", "OAuth login redirect" { + GET "/user" => get_user, "Get current user"; + GET "/config" => get_auth_config, "Get auth config", { + description: "Get the current auth configuration of the server" + }; + GET "/login/{provider}" => oauth_login, "OAuth login", { responses: { 303: () } }; - GET "/login/{provider}/callback" => oauth_callback, "OAuth login callback" { + GET "/login/{provider}/callback" => oauth_callback, "OAuth login callback", { responses: { 303: () } }; - GET, POST "/logout" => logout, "Logout" { + GET, POST "/logout" => logout, "Logout", { responses: { 204: () } }; } diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index 721f93b..9caf904 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -18,10 +18,12 @@ use crate::{ api_routes! { state: AppState, tag: ApiTag::Chat, - POST "/prompt" => prompt, "Prompt", - "Send a single prompt to a provider and get the response"; - POST "/session/{session_id}" => chat_stream, "Chat", - "Send a message in a chat session and stream the response"; + POST "/prompt" => prompt, "Prompt", { + description: "Send a single prompt to a provider and get the response" + }; + POST "/session/{session_id}" => chat_stream, "Chat", { + description: "Send a message in a chat session and stream the response" + }; } #[derive(Debug, Deserialize, JsonSchema)] From 41b70afb177bb3a568927a7c6922e6d29878ff74 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 6 Jul 2026 01:40:48 -0400 Subject: [PATCH 070/111] Update config.rs --- server-new/src/config.rs | 2 ++ 1 file changed, 2 insertions(+) diff --git a/server-new/src/config.rs b/server-new/src/config.rs index 29788ba..b043fac 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -20,6 +20,8 @@ pub fn plugin() -> AdHocPlugin { tracing::info!( log_level = config.server.log_level, base_url = config.server.base_url, + host = %config.server.host, + port = config.server.port, "Config loaded!" ); state.insert(config); From 90c1f243072a10acad677f63f21392ea51a1b367 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Mon, 6 Jul 2026 22:45:22 -0400 Subject: [PATCH 071/111] add provider routes --- server-new/crates/aide-docs-macro/src/lib.rs | 3 +- server-new/src/api/mod.rs | 6 +- server-new/src/api/provider.rs | 71 +++++++++ server-new/src/db/mod.rs | 3 + server-new/src/db/models/provider.rs | 25 +++- server-new/src/db/repositories.rs | 2 + server-new/src/db/repositories/secret.rs | 81 ++++++++++ server-new/src/llm/error.rs | 2 + server-new/src/llm/providers/lorem.rs | 2 +- server-new/src/llm/providers/openai/mod.rs | 4 +- server-new/src/services/provider/mod.rs | 150 ++++++++++++++++--- server-new/src/services/provider/types.rs | 22 +++ 12 files changed, 337 insertions(+), 34 deletions(-) create mode 100644 server-new/src/api/provider.rs create mode 100644 server-new/src/db/repositories/secret.rs create mode 100644 server-new/src/services/provider/types.rs diff --git a/server-new/crates/aide-docs-macro/src/lib.rs b/server-new/crates/aide-docs-macro/src/lib.rs index f0bd4e4..9bd8b52 100644 --- a/server-new/crates/aide-docs-macro/src/lib.rs +++ b/server-new/crates/aide-docs-macro/src/lib.rs @@ -103,11 +103,12 @@ impl RouteMethod { let fn_name = match self.name().as_str() { "GET" => "get_with", "POST" => "post_with", + "PATCH" => "patch_with", "DELETE" => "delete_with", _ => { return Err(syn::Error::new_spanned( &self.ident, - "expected one of GET, POST, DELETE", + "expected one of GET, POST, PATCH, DELETE", )); } }; diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index e99e8b2..43cfa83 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -7,7 +7,7 @@ use aide::{ }; use axum::{Extension, routing::get}; use axum_plugin::AdHocPlugin; -use strum::{AsRefStr, Display, EnumIter, EnumMessage, IntoEnumIterator, IntoStaticStr}; +use strum::{Display, EnumIter, EnumMessage, IntoEnumIterator, IntoStaticStr}; use crate::state::AppState; @@ -15,12 +15,13 @@ pub mod api_key; pub mod auth; pub mod chat; pub mod health; +pub mod provider; const API_BASE: &str = "/api/v1"; const API_AUTH_BASE: &str = "/api/v1/auth"; pub const API_KEY_SCHEME: &str = "ApiKey"; -#[derive(Display, AsRefStr, IntoStaticStr, EnumMessage, EnumIter)] +#[derive(Display, IntoStaticStr, EnumMessage, EnumIter)] enum ApiTag { #[strum(message = "Manage API keys")] ApiKey, @@ -44,6 +45,7 @@ pub fn plugin() -> AdHocPlugin { ) .nest("/chat", chat::routes()) .nest("/health", health::routes()) + .nest("/provider", provider::routes()) .finish_api_with(&mut openapi, build_openapi_doc); let api_routes_with_docs = api_routes diff --git a/server-new/src/api/provider.rs b/server-new/src/api/provider.rs new file mode 100644 index 0000000..f5a1e07 --- /dev/null +++ b/server-new/src/api/provider.rs @@ -0,0 +1,71 @@ +use aide_docs_macro::api_routes; +use axum::{ + Json, + extract::{Path, State}, +}; + +use crate::{ + api::ApiTag, + db::models::ChatRsProvider, + error::AppResult, + extractors::{CurrentUser, Database}, + services::provider::types::{ProviderCreateInput, ProviderUpdateInput}, + state::AppState, +}; + +api_routes! { + state: AppState, + tag: ApiTag::Provider, + GET "/" => list_providers, "List providers"; + POST "/" => create_provider, "Create provider"; + PATCH "/{id}" => update_provider, "Update provider"; + DELETE "/{id}" => delete_provider, "Delete provider"; +} + +async fn list_providers( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, +) -> AppResult>> { + let providers = db.providers().list_by_user_id(&user_id).await?; + Ok(Json(providers)) +} + +async fn create_provider( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, + State(state): State, + Json(input): Json, +) -> AppResult> { + let provider = state + .provider_service() + .create_provider(&mut db, &user_id, &input) + .await?; + Ok(Json(provider)) +} + +async fn update_provider( + CurrentUser { user_id }: CurrentUser, + Path(provider_id): Path, + Database(mut db): Database, + State(state): State, + Json(input): Json, +) -> AppResult> { + let updated_provider = state + .provider_service() + .update_provider(&mut db, &user_id, provider_id, &input) + .await?; + Ok(Json(updated_provider)) +} + +async fn delete_provider( + CurrentUser { user_id }: CurrentUser, + Path(provider_id): Path, + Database(mut db): Database, + State(state): State, +) -> AppResult> { + let deleted_provider = state + .provider_service() + .delete_provider(&mut db, &user_id, provider_id) + .await?; + Ok(Json(deleted_provider)) +} diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs index c1174fc..fbd2933 100644 --- a/server-new/src/db/mod.rs +++ b/server-new/src/db/mod.rs @@ -60,6 +60,9 @@ impl DbService { pub fn providers(&mut self) -> repositories::ProviderRepository<'_> { repositories::ProviderRepository::new(&mut self.cxn) } + pub fn secrets(&mut self) -> repositories::SecretRepository<'_> { + repositories::SecretRepository::new(&mut self.cxn) + } pub fn users(&mut self) -> repositories::UserRepository<'_> { repositories::UserRepository::new(&mut self.cxn) } diff --git a/server-new/src/db/models/provider.rs b/server-new/src/db/models/provider.rs index 16f79ad..90d7ab8 100644 --- a/server-new/src/db/models/provider.rs +++ b/server-new/src/db/models/provider.rs @@ -1,20 +1,21 @@ use chrono::{DateTime, Utc}; use diesel::prelude::*; -use serde::Serialize; +use schemars::JsonSchema; +use serde::{Deserialize, Serialize}; use strum::{EnumString, IntoStaticStr}; use uuid::Uuid; use crate::db::models::ChatRsUser; -#[derive(Identifiable, Associations, Queryable, Selectable, Serialize)] +#[derive(Identifiable, Associations, Queryable, Selectable, Serialize, JsonSchema)] #[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] #[diesel(table_name = super::schema::providers)] pub struct ChatRsProvider { pub id: i32, pub name: String, - // #[schemars(with = "ChatRsProviderType")] + #[schemars(with = "ChatRsProviderType")] pub provider_type: String, - // #[schemars(with = "OpenaiSubtype")] + #[schemars(with = "OpenAISubtype")] pub openai_subtype: Option, #[serde(skip)] pub user_id: Uuid, @@ -29,7 +30,7 @@ pub struct ChatRsProvider { pub struct NewChatRsProvider<'a> { pub name: &'a str, pub provider_type: &'a str, - pub openai_subtype: &'a str, + pub openai_subtype: Option<&'a str>, pub user_id: &'a Uuid, pub base_url: Option<&'a str>, pub default_model: &'a str, @@ -46,17 +47,25 @@ pub struct UpdateChatRsProvider<'a> { } /// The API type of the provider -#[derive(Debug, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr, Deserialize, JsonSchema)] +#[serde(rename_all = "lowercase")] #[strum(serialize_all = "lowercase")] pub enum ChatRsProviderType { - Anthropic, + /// OpenAI or OpenAI-compatible provider OpenAI, + /// Anthropic provider + Anthropic, + /// Ollama provider Ollama, + /// Lorem ipsum provider (for testing) Lorem, } /// The subtype for OpenAI-compatible providers -#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr)] +#[derive( + Debug, Default, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr, Deserialize, JsonSchema, +)] +#[serde(rename_all = "lowercase")] #[strum(serialize_all = "lowercase")] pub enum OpenAISubtype { #[default] diff --git a/server-new/src/db/repositories.rs b/server-new/src/db/repositories.rs index c5b4d12..34df78e 100644 --- a/server-new/src/db/repositories.rs +++ b/server-new/src/db/repositories.rs @@ -1,11 +1,13 @@ mod api_key; mod chat; mod provider; +mod secret; mod session; mod user; pub use api_key::ApiKeyRepository; pub use chat::ChatRepository; pub use provider::ProviderRepository; +pub use secret::SecretRepository; pub use session::SessionRepository; pub use user::UserRepository; diff --git a/server-new/src/db/repositories/secret.rs b/server-new/src/db/repositories/secret.rs new file mode 100644 index 0000000..a4c0de2 --- /dev/null +++ b/server-new/src/db/repositories/secret.rs @@ -0,0 +1,81 @@ +use diesel::prelude::*; +use diesel::result::Error; +use diesel_async::RunQueryDsl; +use uuid::Uuid; + +use crate::db::{ + DbConnection, + models::{ChatRsSecretMeta, NewChatRsSecret, UpdateChatRsSecret}, + schema::secrets, +}; + +pub struct SecretRepository<'a> { + pub db: &'a mut DbConnection, +} + +impl<'a> SecretRepository<'a> { + pub fn new(db: &'a mut DbConnection) -> Self { + SecretRepository { db } + } + + pub async fn find_by_user_id( + &mut self, + user_id: &Uuid, + ) -> Result, Error> { + let keys = secrets::table + .filter(secrets::user_id.eq(user_id)) + .select(ChatRsSecretMeta::as_select()) + .load(self.db) + .await?; + + Ok(keys) + } + + pub async fn create(&mut self, secret: NewChatRsSecret<'_>) -> Result { + let id: Uuid = diesel::insert_into(secrets::table) + .values(secret) + .returning(secrets::id) + .get_result(self.db) + .await?; + + Ok(id) + } + + pub async fn update( + &mut self, + user_id: &Uuid, + secret_id: &Uuid, + data: UpdateChatRsSecret<'_>, + ) -> Result { + let id: Uuid = diesel::update(secrets::table) + .filter(secrets::id.eq(secret_id)) + .filter(secrets::user_id.eq(user_id)) + .set(data) + .returning(secrets::id) + .get_result(self.db) + .await?; + + Ok(id) + } + + pub async fn delete(&mut self, user_id: &Uuid, secret_id: &Uuid) -> Result { + let id: Uuid = diesel::delete(secrets::table) + .filter(secrets::id.eq(secret_id)) + .filter(secrets::user_id.eq(user_id)) + .returning(secrets::id) + .get_result(self.db) + .await?; + + Ok(id) + } + + pub async fn delete_by_user(&mut self, user_id: &Uuid) -> Result, Error> { + let ids: Vec = diesel::delete(secrets::table) + .filter(secrets::user_id.eq(user_id)) + .returning(secrets::id) + .get_results(self.db) + .await?; + + Ok(ids) + } +} diff --git a/server-new/src/llm/error.rs b/server-new/src/llm/error.rs index b68d66b..f555cdd 100644 --- a/server-new/src/llm/error.rs +++ b/server-new/src/llm/error.rs @@ -5,6 +5,8 @@ use crate::services::stream::error::StreamingError; pub enum LlmRequestError { #[error("provider error: {0}")] Provider(String), + #[error("failed to read response: {0}")] + Read(#[from] reqwest::Error), #[error("no content")] NoContent, } diff --git a/server-new/src/llm/providers/lorem.rs b/server-new/src/llm/providers/lorem.rs index 4b2e35f..1d1afc3 100644 --- a/server-new/src/llm/providers/lorem.rs +++ b/server-new/src/llm/providers/lorem.rs @@ -42,7 +42,7 @@ impl Stream for LoremStream { return std::task::Poll::Ready(None); } - match Pin::new(&mut self.interval).poll_tick(cx) { + match self.interval.poll_tick(cx) { std::task::Poll::Ready(_) => { let word = self.words[self.index]; self.index += 1; diff --git a/server-new/src/llm/providers/openai/mod.rs b/server-new/src/llm/providers/openai/mod.rs index 144ff8c..1577ef7 100644 --- a/server-new/src/llm/providers/openai/mod.rs +++ b/server-new/src/llm/providers/openai/mod.rs @@ -148,9 +148,7 @@ impl LlmProvider for OpenAIProvider { &request, ) .await?; - let mut response: OpenAIResponse = response.json().await.map_err(|err| { - LlmRequestError::Provider(format!("Failed to parse response: {err}")) - })?; + let mut response: OpenAIResponse = response.json().await?; let text = response .choices diff --git a/server-new/src/services/provider/mod.rs b/server-new/src/services/provider/mod.rs index e00f324..a82ef87 100644 --- a/server-new/src/services/provider/mod.rs +++ b/server-new/src/services/provider/mod.rs @@ -5,16 +5,26 @@ use uuid::Uuid; use crate::{ db::{ DbService, - models::{ChatRsProvider, ChatRsProviderType, ChatRsSecret, OpenAISubtype}, + models::{ + ChatRsProvider, ChatRsProviderType, ChatRsSecret, NewChatRsProvider, NewChatRsSecret, + OpenAISubtype, UpdateChatRsProvider, UpdateChatRsSecret, + }, }, llm::{ interface::LlmProvider, providers::{LoremProvider, OpenAIProvider, OpenAIProviderConfig}, }, - services::{auth::encryption::Encryptor, provider::error::ProviderError}, + services::{ + auth::encryption::Encryptor, + provider::{ + error::ProviderError, + types::{ProviderCreateInput, ProviderUpdateInput}, + }, + }, }; mod error; +pub mod types; pub struct ProviderService<'r> { encryptor: &'r Encryptor, @@ -29,22 +39,6 @@ impl<'r> ProviderService<'r> { } } - pub async fn get_provider( - &self, - db: &mut DbService, - user_id: &Uuid, - provider_id: i32, - ) -> Result<(ChatRsProvider, ChatRsProviderType, Option), ProviderError> { - let (provider, api_key_secret) = db - .providers() - .find_by_id(user_id, provider_id) - .await? - .ok_or(ProviderError::NotFound)?; - let provider_type = ChatRsProviderType::from_str(&provider.provider_type)?; - - Ok((provider, provider_type, api_key_secret)) - } - pub async fn build_llm_provider( &self, db: &mut DbService, @@ -59,7 +53,6 @@ impl<'r> ProviderService<'r> { .decrypt_string(&secret.ciphertext, &secret.nonce) }) .transpose()?; - let llm_provider: Arc = match provider_type { ChatRsProviderType::Lorem => Arc::new(LoremProvider::new()), ChatRsProviderType::OpenAI => Arc::new(OpenAIProvider::new( @@ -87,4 +80,123 @@ impl<'r> ProviderService<'r> { Ok(llm_provider) } + + pub async fn get_provider( + &self, + db: &mut DbService, + user_id: &Uuid, + provider_id: i32, + ) -> Result<(ChatRsProvider, ChatRsProviderType, Option), ProviderError> { + let (provider, api_key_secret) = db + .providers() + .find_by_id(user_id, provider_id) + .await? + .ok_or(ProviderError::NotFound)?; + let provider_type = ChatRsProviderType::from_str(&provider.provider_type)?; + + Ok((provider, provider_type, api_key_secret)) + } + + pub async fn create_provider( + &self, + db: &mut DbService, + user_id: &Uuid, + input: &ProviderCreateInput, + ) -> Result { + let mut api_key_id: Option = None; + if let Some(plaintext_key) = input.api_key.as_deref() { + let (ciphertext, nonce) = self.encryptor.encrypt_string(plaintext_key)?; + let secret_id = db + .secrets() + .create(NewChatRsSecret { + user_id, + name: &format!("{} API Key", input.name), + ciphertext: &ciphertext, + nonce: &nonce, + }) + .await?; + api_key_id = Some(secret_id); + } + let provider = db + .providers() + .create(NewChatRsProvider { + name: &input.name, + user_id, + provider_type: input.r#type.into(), + openai_subtype: input.openai_type.map(|t| t.into()), + base_url: input.base_url.as_deref(), + default_model: &input.default_model, + api_key_id, + }) + .await?; + + Ok(provider) + } + + pub async fn update_provider( + &self, + db: &mut DbService, + user_id: &Uuid, + provider_id: i32, + input: &ProviderUpdateInput, + ) -> Result { + let (provider, _, secret) = self.get_provider(db, user_id, provider_id).await?; + + let mut secret_id: Option = None; + if let Some(new_api_key) = input.api_key.as_deref() { + let (ciphertext, nonce) = self.encryptor.encrypt_string(new_api_key)?; + secret_id = match secret { + Some(existing_secret) => { + let update_secret = UpdateChatRsSecret { + ciphertext: Some(&ciphertext), + nonce: Some(&nonce), + ..Default::default() + }; + let secret_id = db + .secrets() + .update(user_id, &existing_secret.id, update_secret) + .await?; + Some(secret_id) + } + None => { + let new_secret = NewChatRsSecret { + user_id, + name: &format!("{} API Key", provider.name), + ciphertext: &ciphertext, + nonce: &nonce, + }; + let secret_id = db.secrets().create(new_secret).await?; + Some(secret_id) + } + }; + } + + let update_provider = UpdateChatRsProvider { + api_key_id: secret_id, + name: input.name.as_deref(), + base_url: input.base_url.as_deref(), + default_model: input.default_model.as_deref(), + }; + let updated = db + .providers() + .update(&user_id, provider_id, update_provider) + .await?; + + Ok(updated) + } + + pub async fn delete_provider( + &self, + db: &mut DbService, + user_id: &Uuid, + provider_id: i32, + ) -> Result { + let (_provider, _, api_key_secret) = self.get_provider(db, user_id, provider_id).await?; + if let Some(secret) = api_key_secret { + db.secrets().delete(&user_id, &secret.id).await?; + } + let deleted = db.providers().delete(&user_id, provider_id).await?; + + Ok(deleted) + } } diff --git a/server-new/src/services/provider/types.rs b/server-new/src/services/provider/types.rs new file mode 100644 index 0000000..794957c --- /dev/null +++ b/server-new/src/services/provider/types.rs @@ -0,0 +1,22 @@ +use schemars::JsonSchema; +use serde::Deserialize; + +use crate::db::models::{ChatRsProviderType, OpenAISubtype}; + +#[derive(Deserialize, JsonSchema)] +pub struct ProviderCreateInput { + pub(super) name: String, + pub(super) r#type: ChatRsProviderType, + pub(super) openai_type: Option, + pub(super) base_url: Option, + pub(super) default_model: String, + pub(super) api_key: Option, +} + +#[derive(Deserialize, JsonSchema)] +pub struct ProviderUpdateInput { + pub(super) name: Option, + pub(super) base_url: Option, + pub(super) default_model: Option, + pub(super) api_key: Option, +} From 119e7a14c3974f83cc6ea635b01767d80657fd9b Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 7 Jul 2026 00:17:04 -0400 Subject: [PATCH 072/111] list user's auth sessions --- server-new/src/api/auth.rs | 19 ++++++++--- server-new/src/db/models/session.rs | 6 ++-- server-new/src/db/repositories/api_key.rs | 12 +++---- server-new/src/db/repositories/session.rs | 40 +++++++++++++++-------- server-new/src/services/auth/oauth.rs | 2 +- 5 files changed, 49 insertions(+), 30 deletions(-) diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index df88bee..227f9e4 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -10,7 +10,7 @@ use serde::Deserialize; use crate::{ api::{ApiTag, RoutePrefix}, - db::models::ChatRsUser, + db::models::{ChatRsAuthSession, ChatRsUser}, error::AppResult, extractors::{AppSession, CurrentUser, Database, PublicAuthConfig}, services::auth::oauth::OAuthProviderEnum, @@ -21,6 +21,7 @@ api_routes! { state: AppState, tag: ApiTag::Auth, GET "/user" => get_user, "Get current user"; + GET "/sessions" => list_active_sessions, "List active sessions"; GET "/config" => get_auth_config, "Get auth config", { description: "Get the current auth configuration of the server" }; @@ -44,6 +45,14 @@ async fn get_user( Ok(Json(user)) } +async fn list_active_sessions( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, +) -> AppResult>> { + let sessions = db.auth_sessions().list_active_by_user_id(&user_id).await?; + Ok(Json(sessions)) +} + async fn get_auth_config(auth_config: PublicAuthConfig) -> Json { Json(auth_config) } @@ -79,9 +88,9 @@ async fn oauth_callback( Extension(RoutePrefix(prefix)): Extension, AppSession { session, meta }: AppSession, Database(mut db): Database, - State(app_state): State, + State(state): State, ) -> AppResult { - let oauth = app_state.auth_service().oauth(); + let oauth = state.auth_service().oauth(); let token = oauth .exchange_code( &provider, @@ -94,13 +103,13 @@ async fn oauth_callback( let user = oauth .get_user(&mut db, &provider, &token, maybe_user) .await?; - app_state + state .auth_service() .session() .login(&session, &meta, &user.id) .await?; - Ok(Redirect::to(&app_state.config.server.base_url)) + Ok(Redirect::to(&state.config.server.base_url)) } async fn logout( diff --git a/server-new/src/db/models/session.rs b/server-new/src/db/models/session.rs index 3084efa..fa813e4 100644 --- a/server-new/src/db/models/session.rs +++ b/server-new/src/db/models/session.rs @@ -2,16 +2,18 @@ use std::collections::HashMap; use diesel::{deserialize::FromSqlRow, expression::AsExpression, prelude::*}; use diesel_jsonb_derive::AsJsonb; +use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use uuid::Uuid; use crate::db::{UtcDateTime, models::ChatRsUser}; -#[derive(Identifiable, Associations, Queryable, Selectable)] +#[derive(Identifiable, Associations, Queryable, Selectable, Serialize, JsonSchema)] #[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] #[diesel(table_name = super::schema::auth_sessions)] pub struct ChatRsAuthSession { pub id: Uuid, + #[serde(skip)] pub user_id: Option, pub data: AuthSessionData, pub expires_at: UtcDateTime, @@ -33,6 +35,6 @@ pub struct UpdateChatRsAuthSession { pub expires_at: UtcDateTime, } -#[derive(Debug, Serialize, Deserialize, FromSqlRow, AsExpression, AsJsonb)] +#[derive(Debug, Serialize, Deserialize, FromSqlRow, AsExpression, AsJsonb, JsonSchema)] #[diesel(sql_type = diesel::sql_types::Jsonb)] pub struct AuthSessionData(pub HashMap); diff --git a/server-new/src/db/repositories/api_key.rs b/server-new/src/db/repositories/api_key.rs index ac255e9..239a3a4 100644 --- a/server-new/src/db/repositories/api_key.rs +++ b/server-new/src/db/repositories/api_key.rs @@ -38,13 +38,11 @@ impl<'a> ApiKeyRepository<'a> { } pub async fn create(&mut self, api_key: NewChatRsApiKey<'_>) -> Result { - let id: Uuid = diesel::insert_into(app_api_keys::table) + diesel::insert_into(app_api_keys::table) .values(api_key) .returning(app_api_keys::id) .get_result(self.db) - .await?; - - Ok(id) + .await } pub async fn delete( @@ -62,12 +60,10 @@ impl<'a> ApiKeyRepository<'a> { } pub async fn delete_by_user(&mut self, user_id: &Uuid) -> Result, Error> { - let ids: Vec = diesel::delete(app_api_keys::table) + diesel::delete(app_api_keys::table) .filter(app_api_keys::user_id.eq(user_id)) .returning(app_api_keys::id) .get_results(self.db) - .await?; - - Ok(ids) + .await } } diff --git a/server-new/src/db/repositories/session.rs b/server-new/src/db/repositories/session.rs index a4cfa98..dbe0952 100644 --- a/server-new/src/db/repositories/session.rs +++ b/server-new/src/db/repositories/session.rs @@ -19,6 +19,32 @@ impl<'a> SessionRepository<'a> { Self { db } } + /// Find an active (not expired) session by ID + pub async fn find_active_by_id( + &mut self, + session_id: &Uuid, + ) -> QueryResult> { + auth_sessions::table + .find(session_id) + .filter(auth_sessions::expires_at.gt(diesel::dsl::now)) + .select(ChatRsAuthSession::as_select()) + .first(self.db) + .await + .optional() + } + + pub async fn list_active_by_user_id( + &mut self, + user_id: &Uuid, + ) -> QueryResult> { + auth_sessions::table + .filter(auth_sessions::user_id.eq(user_id)) + .filter(auth_sessions::expires_at.gt(diesel::dsl::now)) + .select(ChatRsAuthSession::as_select()) + .load(self.db) + .await + } + pub async fn create( &mut self, session_id: &Uuid, @@ -54,20 +80,6 @@ impl<'a> SessionRepository<'a> { .await } - /// Find an active (not expired) session by ID - pub async fn find_active_by_id( - &mut self, - session_id: &Uuid, - ) -> QueryResult> { - auth_sessions::table - .find(session_id) - .filter(auth_sessions::expires_at.gt(diesel::dsl::now)) - .select(ChatRsAuthSession::as_select()) - .first(self.db) - .await - .optional() - } - /// Delete a session by ID. Won't return an error if it does not exist. pub async fn delete_by_id(&mut self, session_id: &Uuid) -> QueryResult { diesel::delete(auth_sessions::table.find(session_id)) diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index c26a08d..cbed8a9 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -172,7 +172,7 @@ impl<'a> OAuthService<'a> { } Some(sess) => match db.users().find_by_id(&sess.user_id).await? { Some(user) if oauth_provider.is_user_linked(&user) => { - return Err(AuthError::BadRequest("user already linked to provider")); + return Err(AuthError::BadRequest("already linked to this provider")); } Some(user) => { // Link logged-in user to new provider From 7db282743e83304095f0e122f51d4d9aa34fa73b Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 7 Jul 2026 01:42:46 -0400 Subject: [PATCH 073/111] list models from models.dev --- server-new/src/api/provider.rs | 24 +++++- server-new/src/services/mod.rs | 1 + server-new/src/services/model/error.rs | 26 ++++++ server-new/src/services/model/mod.rs | 108 +++++++++++++++++++++++++ server-new/src/services/model/types.rs | 36 +++++++++ server-new/src/state.rs | 4 + 6 files changed, 198 insertions(+), 1 deletion(-) create mode 100644 server-new/src/services/model/error.rs create mode 100644 server-new/src/services/model/mod.rs create mode 100644 server-new/src/services/model/types.rs diff --git a/server-new/src/api/provider.rs b/server-new/src/api/provider.rs index f5a1e07..fa2b94e 100644 --- a/server-new/src/api/provider.rs +++ b/server-new/src/api/provider.rs @@ -9,7 +9,10 @@ use crate::{ db::models::ChatRsProvider, error::AppResult, extractors::{CurrentUser, Database}, - services::provider::types::{ProviderCreateInput, ProviderUpdateInput}, + services::{ + model::types::LlmModel, + provider::types::{ProviderCreateInput, ProviderUpdateInput}, + }, state::AppState, }; @@ -17,6 +20,7 @@ api_routes! { state: AppState, tag: ApiTag::Provider, GET "/" => list_providers, "List providers"; + GET "/{id}/models" => list_models, "List models"; POST "/" => create_provider, "Create provider"; PATCH "/{id}" => update_provider, "Update provider"; DELETE "/{id}" => delete_provider, "Delete provider"; @@ -30,6 +34,24 @@ async fn list_providers( Ok(Json(providers)) } +async fn list_models( + CurrentUser { user_id }: CurrentUser, + Path(provider_id): Path, + Database(mut db): Database, + State(state): State, +) -> AppResult>> { + let (provider, provider_type, _) = state + .provider_service() + .get_provider(&mut db, &user_id, provider_id) + .await?; + let models = state + .model_service() + .list_models(&provider, &provider_type) + .await?; + + Ok(Json(models)) +} + async fn create_provider( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, diff --git a/server-new/src/services/mod.rs b/server-new/src/services/mod.rs index 1e9f048..95a51c8 100644 --- a/server-new/src/services/mod.rs +++ b/server-new/src/services/mod.rs @@ -1,4 +1,5 @@ pub mod auth; pub mod chat; +pub mod model; pub mod provider; pub mod stream; diff --git a/server-new/src/services/model/error.rs b/server-new/src/services/model/error.rs new file mode 100644 index 0000000..805351f --- /dev/null +++ b/server-new/src/services/model/error.rs @@ -0,0 +1,26 @@ +use crate::error::AppError; + +#[derive(Debug, thiserror::Error)] +pub enum ModelError { + #[error("provider not supported")] + ProviderNotSupported, + #[error("invalid provider type: {0}")] + InvalidProviderType(#[from] strum::ParseError), + #[error("models.dev request error: {0}")] + Request(#[from] reqwest::Error), + #[error("models.dev provider not found: {0}")] + ProviderNotFound(&'static str), + #[error("Redis error: {0}")] + Redis(#[from] fred::prelude::Error), + #[error("Serialization error: {0}")] + Serialization(#[from] serde_json::Error), +} + +impl From for AppError { + fn from(error: ModelError) -> Self { + match error { + ModelError::ProviderNotSupported => Self::bad_request("listing models not supported"), + err => Self::internal(err.into()), + } + } +} diff --git a/server-new/src/services/model/mod.rs b/server-new/src/services/model/mod.rs new file mode 100644 index 0000000..2c20f84 --- /dev/null +++ b/server-new/src/services/model/mod.rs @@ -0,0 +1,108 @@ +use std::{collections::HashMap, str::FromStr}; + +use fred::prelude::{HashesInterface, KeysInterface}; +use serde::Deserialize; +use strum::{AsRefStr, EnumIter, IntoEnumIterator, IntoStaticStr}; + +use crate::{ + db::models::{ChatRsProvider, ChatRsProviderType, OpenAISubtype}, + services::model::{error::ModelError, types::LlmModel}, +}; + +pub mod error; +pub mod types; + +const MODELS_DEV_URL: &str = "https://models.dev/api.json"; +const CACHE_KEY: &str = "rs-chat:models"; +const CACHE_TTL: i64 = 86400; // 1 day in seconds + +/// Service for fetching/listing available LLM models +pub struct ModelService<'r> { + redis: &'r fred::prelude::Pool, + http_client: &'r reqwest::Client, +} + +impl<'r> ModelService<'r> { + pub fn new(redis: &'r fred::prelude::Pool, http_client: &'r reqwest::Client) -> Self { + Self { redis, http_client } + } + + pub async fn list_models( + &self, + provider: &ChatRsProvider, + provider_type: &ChatRsProviderType, + ) -> Result, ModelError> { + let md_provider = match provider_type { + ChatRsProviderType::OpenAI => { + let subtype = provider.openai_subtype.as_deref().unwrap_or_default(); + match OpenAISubtype::from_str(subtype).unwrap_or_default() { + OpenAISubtype::OpenAI => ModelsDevProvider::OpenAI, + OpenAISubtype::OpenRouter => ModelsDevProvider::OpenRouter, + } + } + ChatRsProviderType::Anthropic => ModelsDevProvider::Anthropic, + _ => return Err(ModelError::ProviderNotSupported), + }; + + if let Some(models) = self + .redis + .hget::, _, _>(CACHE_KEY, md_provider.as_ref()) + .await? + .and_then(|models| serde_json::from_str(&models).ok()) + { + Ok(models) + } else { + let mut res: ModelsDevResponse = self + .http_client + .get(MODELS_DEV_URL) + .send() + .await? + .json() + .await?; + + let mut models: Option> = None; + let mut cache: HashMap = HashMap::new(); + for provider in ModelsDevProvider::iter() { + let provider_models: Vec = res + .remove(provider.as_ref()) + .ok_or_else(|| ModelError::ProviderNotFound(provider.into()))? + .models + .into_iter() + .map(|(_, model)| model) + .collect(); + let provider_models_str = serde_json::to_string(&provider_models)?; + cache.insert(provider.as_ref().to_owned(), provider_models_str); + + if md_provider == provider { + models = Some(provider_models); + } + } + + let pipeline = self.redis.next().pipeline(); + let _: () = pipeline.hset(CACHE_KEY, cache).await?; + let _: () = pipeline.expire(CACHE_KEY, CACHE_TTL, None).await?; + let _: () = pipeline.all().await?; + + Ok(models.unwrap_or_default()) + } + } +} + +/// A provider on `models.dev` +#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, IntoStaticStr, AsRefStr, EnumIter)] +#[serde(rename_all = "lowercase")] +#[strum(serialize_all = "lowercase")] +enum ModelsDevProvider { + OpenAI, + OpenRouter, + Anthropic, +} + +/// Map of providers from `models.dev` +type ModelsDevResponse = HashMap; + +/// Provider data on `models.dev` +#[derive(Debug, Deserialize)] +struct ModelsDevProviderData { + models: HashMap, +} diff --git a/server-new/src/services/model/types.rs b/server-new/src/services/model/types.rs new file mode 100644 index 0000000..457ccec --- /dev/null +++ b/server-new/src/services/model/types.rs @@ -0,0 +1,36 @@ +use schemars::JsonSchema; +use serde::{Deserialize, Serialize}; + +/// A model supported by the LLM provider +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] +pub struct LlmModel { + pub id: String, + pub name: String, + pub attachment: Option, + pub reasoning: Option, + pub temperature: Option, + pub tool_call: Option, + pub release_date: Option, + pub knowledge: Option, + pub modalities: Option, + // // Ollama fields + // pub modified_at: Option, + // pub format: Option, + // pub family: Option, +} + +#[derive(Debug, Clone, JsonSchema, Serialize, Deserialize)] +pub struct Modalities { + input: Vec, + output: Vec, +} + +#[derive(Debug, Clone, JsonSchema, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum ModalityType { + Text, + Image, + Audio, + Video, + Pdf, +} diff --git a/server-new/src/state.rs b/server-new/src/state.rs index 55471b9..41523fd 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -10,6 +10,7 @@ use crate::{ services::{ auth::{AuthService, encryption::Encryptor, oauth::OAuthProviderMap}, chat::ChatService, + model::ModelService, provider::ProviderService, stream::tinistream::TinistreamClient, }, @@ -40,6 +41,9 @@ impl AppState { pub fn provider_service(&self) -> ProviderService<'_> { ProviderService::new(&self.encryptor, &self.http_client) } + pub fn model_service(&self) -> ModelService<'_> { + ModelService::new(&self.redis, &self.http_client) + } } impl Deref for AppState { From 292b0f5ddc9318bb214e43aebf362a86718d80b1 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 7 Jul 2026 03:40:26 -0400 Subject: [PATCH 074/111] list models from ollama --- server-new/Cargo.lock | 10 ++--- server-new/Cargo.toml | 4 +- server-new/src/db/mod.rs | 2 + server-new/src/db/models.rs | 2 + server-new/src/db/repositories.rs | 2 + server-new/src/services/model/error.rs | 13 ++---- server-new/src/services/model/mod.rs | 39 +++++++++++----- server-new/src/services/model/providers.rs | 38 ++++++++++++++++ server-new/src/services/model/types.rs | 52 ++++++++++++++++++++-- 9 files changed, 132 insertions(+), 30 deletions(-) create mode 100644 server-new/src/services/model/providers.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index fe8560f..c0733ff 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -133,9 +133,9 @@ dependencies = [ [[package]] name = "anyhow" -version = "1.0.102" +version = "1.0.103" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" +checksum = "2a4385e2e34eb35d6b3efe798b9eb88096925d87726c0798709bf56d9ed84af3" [[package]] name = "arc-swap" @@ -371,7 +371,7 @@ version = "3.9.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6dee98b0db6a962de883bf5d20362dee4d7ca0d12fe39a7c6c73c844e1cd7c1f" dependencies = [ - "darling 0.23.0", + "darling 0.21.3", "ident_case", "prettyplease", "proc-macro2", @@ -3505,9 +3505,9 @@ checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" [[package]] name = "uuid" -version = "1.23.3" +version = "1.23.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "144d6b123cef80b301b8f72a9e2ca4370ddec21950d0a103dd22c437006d2db7" +checksum = "bf80a72845275afea99e7f2b434723d3bc7e38470fcd1c7ed39a599c73319a53" dependencies = [ "getrandom 0.4.3", "js-sys", diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index a60526f..c88cc59 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -13,7 +13,7 @@ aide = { features = ["axum", "axum-json", "axum-query", "macros", "swagger"] } aide-docs-macro = { path = "crates/aide-docs-macro" } -anyhow = "1.0.102" +anyhow = "1.0.103" async-stream = "0.3.6" async-trait = "0.1.89" axum = { version = "0.8.9", features = ["json", "query"] } @@ -100,4 +100,4 @@ tower-sessions-redis-store = { tracing = "0.1.44" tracing-appender = "0.2.5" tracing-subscriber = { version = "0.3.23", features = ["env-filter", "json"] } -uuid = { version = "1.23.3", features = ["serde", "v4"] } +uuid = { version = "1.23.4", features = ["serde", "v4"] } diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs index fbd2933..0cffd78 100644 --- a/server-new/src/db/mod.rs +++ b/server-new/src/db/mod.rs @@ -1,3 +1,5 @@ +//! Database operations + use std::ops::{Deref, DerefMut}; use diesel_async::{ diff --git a/server-new/src/db/models.rs b/server-new/src/db/models.rs index 2b12e99..4e69f44 100644 --- a/server-new/src/db/models.rs +++ b/server-new/src/db/models.rs @@ -1,3 +1,5 @@ +//! Database models + use crate::db::schema; mod api_key; diff --git a/server-new/src/db/repositories.rs b/server-new/src/db/repositories.rs index 34df78e..8db7026 100644 --- a/server-new/src/db/repositories.rs +++ b/server-new/src/db/repositories.rs @@ -1,3 +1,5 @@ +//! Database repositories + mod api_key; mod chat; mod provider; diff --git a/server-new/src/services/model/error.rs b/server-new/src/services/model/error.rs index 805351f..b51ecdd 100644 --- a/server-new/src/services/model/error.rs +++ b/server-new/src/services/model/error.rs @@ -2,14 +2,10 @@ use crate::error::AppError; #[derive(Debug, thiserror::Error)] pub enum ModelError { - #[error("provider not supported")] - ProviderNotSupported, - #[error("invalid provider type: {0}")] - InvalidProviderType(#[from] strum::ParseError), - #[error("models.dev request error: {0}")] + #[error("request error: {0}")] Request(#[from] reqwest::Error), #[error("models.dev provider not found: {0}")] - ProviderNotFound(&'static str), + ModelsDevProviderNotFound(&'static str), #[error("Redis error: {0}")] Redis(#[from] fred::prelude::Error), #[error("Serialization error: {0}")] @@ -18,9 +14,6 @@ pub enum ModelError { impl From for AppError { fn from(error: ModelError) -> Self { - match error { - ModelError::ProviderNotSupported => Self::bad_request("listing models not supported"), - err => Self::internal(err.into()), - } + Self::internal(error.into()) } } diff --git a/server-new/src/services/model/mod.rs b/server-new/src/services/model/mod.rs index 2c20f84..f28326d 100644 --- a/server-new/src/services/model/mod.rs +++ b/server-new/src/services/model/mod.rs @@ -10,6 +10,7 @@ use crate::{ }; pub mod error; +mod providers; pub mod types; const MODELS_DEV_URL: &str = "https://models.dev/api.json"; @@ -32,18 +33,34 @@ impl<'r> ModelService<'r> { provider: &ChatRsProvider, provider_type: &ChatRsProviderType, ) -> Result, ModelError> { - let md_provider = match provider_type { + match provider_type { ChatRsProviderType::OpenAI => { let subtype = provider.openai_subtype.as_deref().unwrap_or_default(); - match OpenAISubtype::from_str(subtype).unwrap_or_default() { + let md_provider = match OpenAISubtype::from_str(subtype).unwrap_or_default() { OpenAISubtype::OpenAI => ModelsDevProvider::OpenAI, OpenAISubtype::OpenRouter => ModelsDevProvider::OpenRouter, - } + }; + self.fetch_models_dev(md_provider).await + } + ChatRsProviderType::Anthropic => { + self.fetch_models_dev(ModelsDevProvider::Anthropic).await } - ChatRsProviderType::Anthropic => ModelsDevProvider::Anthropic, - _ => return Err(ModelError::ProviderNotSupported), - }; + ChatRsProviderType::Ollama => { + providers::ollama_models(self.http_client, provider.base_url.as_deref()).await + } + ChatRsProviderType::Lorem => Ok(vec![LlmModel { + id: String::from("lorem"), + name: String::from("lorem"), + ..Default::default() + }]), + } + } + /// Fetch detailed model list from `models.dev` with caching + async fn fetch_models_dev( + &self, + md_provider: ModelsDevProvider, + ) -> Result, ModelError> { if let Some(models) = self .redis .hget::, _, _>(CACHE_KEY, md_provider.as_ref()) @@ -57,15 +74,17 @@ impl<'r> ModelService<'r> { .get(MODELS_DEV_URL) .send() .await? + .error_for_status()? .json() .await?; let mut models: Option> = None; let mut cache: HashMap = HashMap::new(); + for provider in ModelsDevProvider::iter() { let provider_models: Vec = res .remove(provider.as_ref()) - .ok_or_else(|| ModelError::ProviderNotFound(provider.into()))? + .ok_or_else(|| ModelError::ModelsDevProviderNotFound(provider.into()))? .models .into_iter() .map(|(_, model)| model) @@ -88,7 +107,7 @@ impl<'r> ModelService<'r> { } } -/// A provider on `models.dev` +/// Represents a provider from `models.dev` #[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, IntoStaticStr, AsRefStr, EnumIter)] #[serde(rename_all = "lowercase")] #[strum(serialize_all = "lowercase")] @@ -98,10 +117,10 @@ enum ModelsDevProvider { Anthropic, } -/// Map of providers from `models.dev` +/// Main response from `models.dev` - map of providers type ModelsDevResponse = HashMap; -/// Provider data on `models.dev` +/// Provider data from `models.dev` #[derive(Debug, Deserialize)] struct ModelsDevProviderData { models: HashMap, diff --git a/server-new/src/services/model/providers.rs b/server-new/src/services/model/providers.rs new file mode 100644 index 0000000..4e9cfbd --- /dev/null +++ b/server-new/src/services/model/providers.rs @@ -0,0 +1,38 @@ +use super::{ + error::ModelError, + types::{LlmModel, OllamaModelsResponse}, +}; + +pub(super) async fn ollama_models( + client: &reqwest::Client, + base_url: Option<&str>, +) -> Result, ModelError> { + const DEFAULT_BASE_URL: &str = "http://localhost:11434"; + const MODELS_API_PATH: &str = "/api/tags"; + + let models_url = format!("{}{MODELS_API_PATH}", base_url.unwrap_or(DEFAULT_BASE_URL)); + let response: OllamaModelsResponse = client + .get(models_url) + .send() + .await? + .error_for_status()? + .json() + .await?; + let models = response + .models + .into_iter() + .map(|model| LlmModel { + id: model.name.clone(), + name: model.name, + temperature: Some(true), + modified_at: Some(model.modified_at), + format: Some(model.details.format), + family: Some(model.details.family), + tool_call: Some(model.capabilities.iter().any(|c| c.is_tools())), + reasoning: Some(model.capabilities.iter().any(|c| c.is_thinking())), + ..Default::default() + }) + .collect(); + + Ok(models) +} diff --git a/server-new/src/services/model/types.rs b/server-new/src/services/model/types.rs index 457ccec..9d8dbd1 100644 --- a/server-new/src/services/model/types.rs +++ b/server-new/src/services/model/types.rs @@ -1,7 +1,9 @@ use schemars::JsonSchema; use serde::{Deserialize, Serialize}; +use serde_with::skip_serializing_none; /// A model supported by the LLM provider +#[skip_serializing_none] #[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] pub struct LlmModel { pub id: String, @@ -14,9 +16,9 @@ pub struct LlmModel { pub knowledge: Option, pub modalities: Option, // // Ollama fields - // pub modified_at: Option, - // pub format: Option, - // pub family: Option, + pub modified_at: Option, + pub format: Option, + pub family: Option, } #[derive(Debug, Clone, JsonSchema, Serialize, Deserialize)] @@ -34,3 +36,47 @@ pub enum ModalityType { Video, Pdf, } + +/// Ollama models list response +#[derive(Debug, Deserialize)] +pub struct OllamaModelsResponse { + pub models: Vec, +} + +/// Ollama model information +#[derive(Debug, Deserialize)] +pub struct OllamaModelInfo { + pub name: String, + // pub model: String, + pub modified_at: String, + // pub size: u64, + // pub digest: String, + pub details: OllamaModelDetails, + #[serde(default)] + pub capabilities: Vec, +} + +/// Ollama model details +#[derive(Debug, Deserialize)] +pub struct OllamaModelDetails { + // #[serde(default)] + // pub parent_model: String, + pub format: String, + pub family: String, + // #[serde(default)] + // pub families: Vec, + // pub parameter_size: String, + // #[serde(default)] + // pub quantization_level: Option, +} + +#[derive(Debug, Deserialize, strum::EnumIs)] +#[serde(rename_all = "lowercase")] +pub enum OllamaCapabilities { + Completion, + Tools, + Vision, + Thinking, + #[serde(other)] + Unknown, +} From 9b315d117e7ec6d8c756fc17a4a332d122388910 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 7 Jul 2026 03:52:35 -0400 Subject: [PATCH 075/111] ollama model tweaks --- server-new/src/services/model/providers.rs | 2 +- server-new/src/services/model/types.rs | 3 ++- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/server-new/src/services/model/providers.rs b/server-new/src/services/model/providers.rs index 4e9cfbd..f1b7243 100644 --- a/server-new/src/services/model/providers.rs +++ b/server-new/src/services/model/providers.rs @@ -22,7 +22,7 @@ pub(super) async fn ollama_models( .models .into_iter() .map(|model| LlmModel { - id: model.name.clone(), + id: model.model, name: model.name, temperature: Some(true), modified_at: Some(model.modified_at), diff --git a/server-new/src/services/model/types.rs b/server-new/src/services/model/types.rs index 9d8dbd1..04d01fc 100644 --- a/server-new/src/services/model/types.rs +++ b/server-new/src/services/model/types.rs @@ -6,6 +6,7 @@ use serde_with::skip_serializing_none; #[skip_serializing_none] #[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] pub struct LlmModel { + /// The model ID to use in a chat / prompt request pub id: String, pub name: String, pub attachment: Option, @@ -47,7 +48,7 @@ pub struct OllamaModelsResponse { #[derive(Debug, Deserialize)] pub struct OllamaModelInfo { pub name: String, - // pub model: String, + pub model: String, pub modified_at: String, // pub size: u64, // pub digest: String, From 78990996e6a3b91f86603bfa6aa88c53720ab546 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 14 Jul 2026 02:22:01 -0400 Subject: [PATCH 076/111] update tinistream and oauth crates --- server-new/Cargo.lock | 15 +++++----- server-new/Cargo.toml | 4 +-- server-new/src/services/auth/oauth.rs | 7 +++-- server-new/src/services/stream/tinistream.rs | 30 +++----------------- 4 files changed, 18 insertions(+), 38 deletions(-) diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index c0733ff..2da41f3 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -371,7 +371,7 @@ version = "3.9.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6dee98b0db6a962de883bf5d20362dee4d7ca0d12fe39a7c6c73c844e1cd7c1f" dependencies = [ - "darling 0.21.3", + "darling 0.23.0", "ident_case", "prettyplease", "proc-macro2", @@ -1987,9 +1987,9 @@ dependencies = [ [[package]] name = "progenitor-client" -version = "0.12.0" +version = "0.14.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ffab7b358944dba033a7b324e7558e66e6bcb1fb4705cf57f26fd5092bcae630" +checksum = "4e8a874cf25a33cac7a01b9c1de87bcfbc8aea93f3156d09dcc3bee516a78926" dependencies = [ "bytes", "futures-core", @@ -2722,9 +2722,9 @@ checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" [[package]] name = "simple-oauth" -version = "0.1.0-beta" +version = "0.1.0-beta.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "46d3a9a3b349db65f2bb743ed1b29ea4e866b25eee9fea0fbfecd805f676f5c0" +checksum = "5b0b2fb63e098c4e51b55768a3bf1b29c0180de4dda801cdd9b568d92967787f" dependencies = [ "bon", "oauth2", @@ -2732,7 +2732,6 @@ dependencies = [ "reqwest", "serde", "serde_json", - "subtle", "thiserror 2.0.18", ] @@ -2942,8 +2941,8 @@ dependencies = [ [[package]] name = "tinistream-client" -version = "0.1.10" -source = "git+https://github.com/fa-sharp/tinistream?rev=f25144c#f25144c1bdbee827d6033606b94a8aa1ae6eb5a7" +version = "0.2.0" +source = "git+https://github.com/fa-sharp/tinistream?rev=015d307#015d3076e64b55fbc6f1577716ede64d6e66597e" dependencies = [ "bytes", "futures-core", diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index c88cc59..fa84521 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -65,7 +65,7 @@ serde_with = { default-features = false, features = ["macros"] } -simple-oauth = { version = "0.1.0-beta", features = ["default-tls"] } +simple-oauth = { version = "0.1.0-beta.1", features = ["default-tls"] } strum = { version = "0.28.0", default-features = false, @@ -74,7 +74,7 @@ strum = { thiserror = "2.0.18" tinistream-client = { git = "https://github.com/fa-sharp/tinistream", - rev = "f25144c" + rev = "015d307" } tokio = { version = "1.52.3", diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index cbed8a9..8a157ce 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -130,14 +130,17 @@ impl<'a> OAuthService<'a> { .await? .ok_or(AuthError::Unauthorized("missing PKCE in session"))?; + // Verify state + if initial_state != state { + return Err(AuthError::Unauthorized("state mismatch")); + } + // Exchange code for token let (oauth_client, _) = self.oauth_provider(provider)?; let response = oauth_client .exchange_code() .redirect_url(self.get_redirect_url(callback_path)) .code(code) - .state(state) - .initial_state(&initial_state) .pkce_verifier(pkce_verifier) .build() .await?; diff --git a/server-new/src/services/stream/tinistream.rs b/server-new/src/services/stream/tinistream.rs index 72706cd..16abaef 100644 --- a/server-new/src/services/stream/tinistream.rs +++ b/server-new/src/services/stream/tinistream.rs @@ -1,7 +1,7 @@ //! Client for `tinistream` to handle streaming responses use reqwest_websocket::{Upgrade, WebSocket}; -use tinistream_client::{Client, ClientEventsExt, ClientInfo, ClientStreamExt, Error, types::*}; +use tinistream_client::{Client, ClientInfo, ClientStreamExt, Error, types::*}; /// A client for interacting with the `tinistream` API. #[derive(Debug, Clone)] @@ -16,7 +16,6 @@ pub type TiniResult = Result; #[error("{status} {message}")] pub struct TiniError { pub status: u16, - pub code: String, pub message: String, } @@ -88,24 +87,6 @@ impl TinistreamClient { res.into_websocket().await } - pub async fn stream_add( - &self, - key: &str, - events: Vec, - ) -> TiniResult> { - let events = events - .into_iter() - .map(|event| event.try_into()) - .collect::, _>>()?; - let res = self - .client - .add_events() - .body(AddEventsRequest::builder().key(key).events(events)) - .send() - .await?; - Ok(res.into_inner().ids) - } - /// Cancel a stream pub async fn stream_cancel(&self, key: &str) -> TiniResult { let res = self @@ -129,21 +110,19 @@ impl TinistreamClient { } } -impl From> for TiniError { - fn from(value: Error) -> Self { +impl From> for TiniError { + fn from(value: Error) -> Self { match value { Error::ErrorResponse(res) => { let status = res.status().as_u16(); let res = res.into_inner(); TiniError { status, - code: res.code, - message: res.message, + message: res.error.message, } } res => TiniError { status: res.status().map_or(500, |s| s.as_u16()), - code: "unexpected".to_owned(), message: res.to_string(), }, } @@ -154,7 +133,6 @@ impl From for TiniError { fn from(value: error::ConversionError) -> Self { TiniError { status: 400, - code: "invalid_event".to_owned(), message: value.to_string(), } } From 502f00e5b660b07dd31bdcaaf3aa298bdc580a26 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 14 Jul 2026 02:28:48 -0400 Subject: [PATCH 077/111] update aide macros crate --- server-new/Cargo.lock | 21 +- server-new/Cargo.toml | 5 +- server-new/crates/aide-docs-macro/Cargo.toml | 14 - server-new/crates/aide-docs-macro/src/lib.rs | 436 ------------------- server-new/src/api/api_key.rs | 4 +- server-new/src/api/auth.rs | 4 +- server-new/src/api/chat.rs | 4 +- server-new/src/api/health.rs | 2 +- server-new/src/api/provider.rs | 4 +- 9 files changed, 24 insertions(+), 470 deletions(-) delete mode 100644 server-new/crates/aide-docs-macro/Cargo.toml delete mode 100644 server-new/crates/aide-docs-macro/src/lib.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index 2da41f3..d239366 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -102,15 +102,6 @@ dependencies = [ "tracing", ] -[[package]] -name = "aide-docs-macro" -version = "0.1.0" -dependencies = [ - "proc-macro2", - "quote", - "syn", -] - [[package]] name = "aide-macros" version = "0.16.0-alpha.4" @@ -271,6 +262,16 @@ dependencies = [ "tracing", ] +[[package]] +name = "axum-aide-macros" +version = "0.1.0" +source = "git+https://git.fasharp.io/fa-sharp/axum-aide-macros?rev=5b00e645df#5b00e645dfec6a0a76cdfa0e9a0c9e050003faec" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "axum-core" version = "0.5.6" @@ -2302,11 +2303,11 @@ version = "0.1.0" dependencies = [ "aes-gcm 0.11.0", "aide", - "aide-docs-macro", "anyhow", "async-stream", "async-trait", "axum", + "axum-aide-macros", "axum-helmet", "axum-plugin", "chrono", diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index fa84521..41d764a 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -12,11 +12,14 @@ aide = { rev = "7246c20", features = ["axum", "axum-json", "axum-query", "macros", "swagger"] } -aide-docs-macro = { path = "crates/aide-docs-macro" } anyhow = "1.0.103" async-stream = "0.3.6" async-trait = "0.1.89" axum = { version = "0.8.9", features = ["json", "query"] } +axum-aide-macros = { + git = "https://git.fasharp.io/fa-sharp/axum-aide-macros", + rev = "5b00e645df" +} axum-helmet = "1.0.2" axum-plugin = { git = "https://git.fasharp.io/fa-sharp/axum-plugin", diff --git a/server-new/crates/aide-docs-macro/Cargo.toml b/server-new/crates/aide-docs-macro/Cargo.toml deleted file mode 100644 index d01c3ff..0000000 --- a/server-new/crates/aide-docs-macro/Cargo.toml +++ /dev/null @@ -1,14 +0,0 @@ -[package] -name = "aide-docs-macro" -version = "0.1.0" -edition = "2024" -description = "Convenience macros for aide API docs" -publish = false - -[lib] -proc-macro = true - -[dependencies] -proc-macro2 = "1" -quote = "1" -syn = { version = "2", features = ["full", "parsing"] } diff --git a/server-new/crates/aide-docs-macro/src/lib.rs b/server-new/crates/aide-docs-macro/src/lib.rs deleted file mode 100644 index 9bd8b52..0000000 --- a/server-new/crates/aide-docs-macro/src/lib.rs +++ /dev/null @@ -1,436 +0,0 @@ -use proc_macro::TokenStream; -use quote::{format_ident, quote}; -use syn::{ - Expr, Ident, ItemFn, LitInt, LitStr, Result, Token, Type, braced, parse::Parse, - parse::ParseStream, parse_macro_input, -}; - -struct DocsArgs { - summary: LitStr, - description: Option, -} - -impl Parse for DocsArgs { - fn parse(input: ParseStream<'_>) -> Result { - let summary = input.parse()?; - let description = if input.peek(Token![,]) { - input.parse::()?; - Some(input.parse()?) - } else { - None - }; - - if !input.is_empty() { - input.parse::()?; - } - - Ok(Self { - summary, - description, - }) - } -} - -/// Convenience macro for generating API docs for the route handler. Generates a function -/// called `_docs` that can be passed as the transform function -/// to aide's `get_with`, `post_with`, etc. -/// -/// # Syntax -/// `#[handler_docs("" (, "")]` -#[proc_macro_attribute] -pub fn handler_docs(args: TokenStream, input: TokenStream) -> TokenStream { - let args = parse_macro_input!(args as DocsArgs); - let handler = parse_macro_input!(input as ItemFn); - let handler_name = &handler.sig.ident; - let docs_name = format_ident!("{}_docs", handler_name); - let operation_id = handler_name.to_string(); - let summary = args.summary; - - let description = args.description.map(|description| { - quote! { - .description(#description) - } - }); - - quote! { - fn #docs_name( - op: ::aide::transform::TransformOperation, - ) -> ::aide::transform::TransformOperation { - op.id(#operation_id) - .summary(#summary) - #description - } - - #handler - } - .into() -} - -struct ApiRoutes { - state: Type, - tag: Option, - routes: Vec, -} - -struct ApiRoute { - methods: Vec, - path: LitStr, - handler: Ident, - summary: Option, - description: Option, - responses: Vec, -} - -struct RouteResponse { - status: LitInt, - ty: Type, -} - -struct RouteMethod { - ident: Ident, -} - -impl RouteMethod { - fn name(&self) -> String { - self.ident.to_string() - } - - fn operation_prefix(&self) -> String { - self.name().to_lowercase() - } - - fn route_fn(&self) -> Result { - let fn_name = match self.name().as_str() { - "GET" => "get_with", - "POST" => "post_with", - "PATCH" => "patch_with", - "DELETE" => "delete_with", - _ => { - return Err(syn::Error::new_spanned( - &self.ident, - "expected one of GET, POST, PATCH, DELETE", - )); - } - }; - - Ok(format_ident!("{fn_name}")) - } -} - -impl Parse for ApiRoutes { - fn parse(input: ParseStream<'_>) -> Result { - parse_label(input, "state")?; - input.parse::()?; - let state = input.parse()?; - input.parse::()?; - - let tag = if next_label_is(input, "tag") { - parse_label(input, "tag")?; - input.parse::()?; - let tag = input.parse()?; - input.parse::()?; - Some(tag) - } else { - None - }; - - let mut routes = Vec::new(); - while !input.is_empty() { - routes.push(input.parse()?); - } - - Ok(Self { state, tag, routes }) - } -} - -impl Parse for ApiRoute { - fn parse(input: ParseStream<'_>) -> Result { - let mut methods = vec![RouteMethod { - ident: input.parse()?, - }]; - - while input.peek(Token![,]) { - let fork = input.fork(); - fork.parse::()?; - if fork.peek(Ident) { - input.parse::()?; - methods.push(RouteMethod { - ident: input.parse()?, - }); - } else { - break; - } - } - - let path = input.parse()?; - input.parse::]>()?; - let handler = input.parse()?; - input.parse::()?; - - let summary = if input.peek(LitStr) { - Some(input.parse()?) - } else { - None - }; - - let options = if input.peek(Token![,]) { - input.parse::()?; - if !input.peek(syn::token::Brace) { - return Err(input.error("expected route options block after summary comma")); - } - input.parse()? - } else if input.peek(syn::token::Brace) { - input.parse()? - } else { - RouteOptions::default() - }; - - input.parse::()?; - - if summary.is_none() && options.is_empty() { - return Err(input.error("expected summary string or route options block")); - } - - Ok(Self { - methods, - path, - handler, - summary, - description: options.description, - responses: options.responses, - }) - } -} - -#[derive(Default)] -struct RouteOptions { - description: Option, - responses: Vec, -} - -impl RouteOptions { - fn is_empty(&self) -> bool { - self.description.is_none() && self.responses.is_empty() - } -} - -impl Parse for RouteOptions { - fn parse(input: ParseStream<'_>) -> Result { - let content; - braced!(content in input); - - let mut options = RouteOptions::default(); - while !content.is_empty() { - let label: Ident = content.parse()?; - content.parse::()?; - - match label.to_string().as_str() { - "description" => { - if options.description.is_some() { - return Err(syn::Error::new_spanned( - label, - "`description` can only be provided once", - )); - } - - options.description = Some(content.parse()?); - } - "responses" => { - if !options.responses.is_empty() { - return Err(syn::Error::new_spanned( - label, - "`responses` can only be provided once", - )); - } - - let responses; - braced!(responses in content); - - while !responses.is_empty() { - options.responses.push(responses.parse()?); - if responses.peek(Token![,]) { - responses.parse::()?; - } - } - } - _ => { - return Err(syn::Error::new_spanned( - label, - "expected `description` or `responses`", - )); - } - } - - if content.peek(Token![,]) { - content.parse::()?; - } - } - - Ok(options) - } -} - -impl Parse for RouteResponse { - fn parse(input: ParseStream<'_>) -> Result { - let status = input.parse()?; - input.parse::()?; - let ty = input.parse()?; - - Ok(Self { status, ty }) - } -} - -fn parse_label(input: ParseStream<'_>, expected: &str) -> Result<()> { - let label: Ident = input.parse()?; - if label == expected { - Ok(()) - } else { - Err(syn::Error::new_spanned( - label, - format!("expected `{expected}`"), - )) - } -} - -fn next_label_is(input: ParseStream<'_>, expected: &str) -> bool { - let fork = input.fork(); - fork.parse::().is_ok_and(|label| label == expected) && fork.peek(Token![:]) -} - -/// Generate a `routes()` function with attached API docs -/// -/// # Syntax -/// ```ignore -/// api_routes! { -/// state: AppState, -/// tag: ApiTag::Auth, // optional -/// GET "/user" => get_user, "Get user", { -/// description: "Get the current user" -/// }; -/// GET, POST "/logout" => logout, "Logout", { -/// responses: { 204: () } -/// }; -/// POST "/sessions" => create_session, { -/// description: "Create a new session", -/// responses: { 201: Session } -/// }; -/// } -/// ``` -#[proc_macro] -pub fn api_routes(input: TokenStream) -> TokenStream { - let api_routes = parse_macro_input!(input as ApiRoutes); - let state = api_routes.state; - - let docs_functions = api_routes.routes.iter().flat_map(|route| { - let summary = route.summary.as_ref().map(|summary| { - quote! { - .summary(#summary) - } - }); - let description = route.description.as_ref().map(|description| { - quote! { - .description(#description) - } - }); - let responses = route - .responses - .iter() - .map(|response| { - let status = &response.status; - let ty = &response.ty; - - quote! { - .response::<#status, #ty>() - } - }) - .collect::>(); - let multiple_methods = route.methods.len() > 1; - - route.methods.iter().map(move |method| { - let docs_name = docs_name(route, method, multiple_methods); - let operation_id = operation_id(route, method, multiple_methods); - - quote! { - fn #docs_name( - op: ::aide::transform::TransformOperation, - ) -> ::aide::transform::TransformOperation { - op.id(#operation_id) - #summary - #description - #(#responses)* - } - } - }) - }); - - let route_calls = api_routes.routes.iter().map(|route| { - let path = &route.path; - let handler = &route.handler; - let multiple_methods = route.methods.len() > 1; - - let mut methods = route - .methods - .iter() - .map(|method| { - let route_fn = method.route_fn()?; - let docs_name = docs_name(route, method, multiple_methods); - - Ok((route_fn, docs_name)) - }) - .collect::>>()?; - let (first_method, first_docs_name) = methods.remove(0); - - let additional_methods = methods.into_iter().map(|(method, docs_name)| { - quote! { - .#method(#handler, #docs_name) - } - }); - - Ok(quote! { - .api_route( - #path, - ::aide::axum::routing::#first_method(#handler, #first_docs_name) - #(#additional_methods)* - ) - }) - }); - - let route_calls = match route_calls.collect::>>() { - Ok(route_calls) => route_calls, - Err(error) => return error.into_compile_error().into(), - }; - - let tag = api_routes.tag.map(|tag| { - quote! { - .with_path_items(|op| op.tag(#tag.into())) - } - }); - - quote! { - #(#docs_functions)* - - pub fn routes() -> ::aide::axum::ApiRouter<#state> { - ::aide::axum::ApiRouter::new() - #(#route_calls)* - #tag - } - } - .into() -} - -fn docs_name(route: &ApiRoute, method: &RouteMethod, multiple_methods: bool) -> Ident { - if multiple_methods { - let method = method.operation_prefix(); - format_ident!("{}_{}_docs", method, route.handler) - } else { - format_ident!("{}_docs", route.handler) - } -} - -fn operation_id(route: &ApiRoute, method: &RouteMethod, multiple_methods: bool) -> String { - if multiple_methods { - format!("{}_{}", method.operation_prefix(), route.handler) - } else { - route.handler.to_string() - } -} diff --git a/server-new/src/api/api_key.rs b/server-new/src/api/api_key.rs index a116932..240934c 100644 --- a/server-new/src/api/api_key.rs +++ b/server-new/src/api/api_key.rs @@ -1,9 +1,9 @@ -use aide_docs_macro::api_routes; use axum::{ Json, extract::{Path, State}, http::StatusCode, }; +use axum_aide_macros::api_routes; use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use uuid::Uuid; @@ -18,7 +18,7 @@ use crate::{ api_routes! { state: AppState, - tag: ApiTag::ApiKey, + tag: ApiTag::ApiKey.into(), GET "/" => list_api_keys, "List API keys"; POST "/" => create_api_key, "Create API key"; DELETE "/{id}" => delete_api_key, "Delete API key"; diff --git a/server-new/src/api/auth.rs b/server-new/src/api/auth.rs index 227f9e4..2d2d8e1 100644 --- a/server-new/src/api/auth.rs +++ b/server-new/src/api/auth.rs @@ -1,10 +1,10 @@ -use aide_docs_macro::api_routes; use axum::{ Extension, Json, extract::{Path, Query, State}, http::StatusCode, response::Redirect, }; +use axum_aide_macros::api_routes; use schemars::JsonSchema; use serde::Deserialize; @@ -19,7 +19,7 @@ use crate::{ api_routes! { state: AppState, - tag: ApiTag::Auth, + tag: ApiTag::Auth.into(), GET "/user" => get_user, "Get current user"; GET "/sessions" => list_active_sessions, "List active sessions"; GET "/config" => get_auth_config, "Get auth config", { diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index 9caf904..9578e4e 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -1,8 +1,8 @@ -use aide_docs_macro::api_routes; use axum::{ Json, extract::{Path, State}, }; +use axum_aide_macros::api_routes; use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use uuid::Uuid; @@ -17,7 +17,7 @@ use crate::{ api_routes! { state: AppState, - tag: ApiTag::Chat, + tag: ApiTag::Chat.into(), POST "/prompt" => prompt, "Prompt", { description: "Send a single prompt to a provider and get the response" }; diff --git a/server-new/src/api/health.rs b/server-new/src/api/health.rs index 669c571..a657c02 100644 --- a/server-new/src/api/health.rs +++ b/server-new/src/api/health.rs @@ -1,5 +1,5 @@ use aide::axum::{ApiRouter, routing::get_with}; -use aide_docs_macro::handler_docs; +use axum_aide_macros::handler_docs; use crate::state::AppState; diff --git a/server-new/src/api/provider.rs b/server-new/src/api/provider.rs index fa2b94e..3ade43a 100644 --- a/server-new/src/api/provider.rs +++ b/server-new/src/api/provider.rs @@ -1,8 +1,8 @@ -use aide_docs_macro::api_routes; use axum::{ Json, extract::{Path, State}, }; +use axum_aide_macros::api_routes; use crate::{ api::ApiTag, @@ -18,7 +18,7 @@ use crate::{ api_routes! { state: AppState, - tag: ApiTag::Provider, + tag: ApiTag::Provider.into(), GET "/" => list_providers, "List providers"; GET "/{id}/models" => list_models, "List models"; POST "/" => create_provider, "Create provider"; From 89879449e755247725b5003cfab91092e3eafee4 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 14 Jul 2026 03:24:12 -0400 Subject: [PATCH 078/111] update axum_plugin crate --- docker-compose.yml | 8 ++-- server-new/Cargo.lock | 9 ++-- server-new/Cargo.toml | 4 +- server-new/src/api/mod.rs | 12 ++--- server-new/src/config.rs | 48 ++++++-------------- server-new/src/lib.rs | 18 ++++++-- server-new/src/plugins/auth.rs | 38 +++++++--------- server-new/src/plugins/clients.rs | 20 ++++---- server-new/src/plugins/database.rs | 28 +++++------- server-new/src/plugins/logging.rs | 11 ++--- server-new/src/plugins/mod.rs | 7 +++ server-new/src/plugins/redis.rs | 36 +++++++-------- server-new/src/plugins/security.rs | 11 ++--- server-new/src/services/stream/tinistream.rs | 11 +++++ server-new/src/state.rs | 2 +- 15 files changed, 124 insertions(+), 139 deletions(-) diff --git a/docker-compose.yml b/docker-compose.yml index 687cc5e..2974726 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -20,18 +20,16 @@ services: - redis_data:/data stream: - image: ghcr.io/fa-sharp/tinistream:0.1.10 - platform: linux/amd64 + image: tinistream container_name: tinistream ports: - "8081:8081" environment: STREAMER_PORT: 8081 STREAMER_API_KEY: dev-streamer-api-key - STREAMER_SERVER_ADDRESS: http://localhost:8081 + STREAMER_BASE_URL: http://localhost:8081 STREAMER_REDIS_URL: redis://redis:6379 - STREAMER_TTL: 360 - env_file: server/.env + env_file: server-new/.env depends_on: - redis diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index d239366..c791db6 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -306,20 +306,22 @@ dependencies = [ [[package]] name = "axum-plugin" -version = "0.2.0" -source = "git+https://git.fasharp.io/fa-sharp/axum-plugin?rev=9f72278b3c#9f72278b3c8ae57897ac04af54a6643bb9e5a25c" +version = "0.3.0" +source = "git+https://git.fasharp.io/fa-sharp/axum-plugin?rev=be17dc9aec#be17dc9aec0f1138131052924befe564834798f9" dependencies = [ "anyhow", "axum", "axum-plugin-macros", + "figment", "futures", + "serde", "type-map", ] [[package]] name = "axum-plugin-macros" version = "0.1.0" -source = "git+https://git.fasharp.io/fa-sharp/axum-plugin?rev=9f72278b3c#9f72278b3c8ae57897ac04af54a6643bb9e5a25c" +source = "git+https://git.fasharp.io/fa-sharp/axum-plugin?rev=be17dc9aec#be17dc9aec0f1138131052924befe564834798f9" dependencies = [ "proc-macro2", "quote", @@ -2317,7 +2319,6 @@ dependencies = [ "diesel-jsonb-derive", "diesel_migrations", "dotenvy", - "figment", "fred", "futures", "hex", diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 41d764a..ba931b5 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -23,7 +23,8 @@ axum-aide-macros = { axum-helmet = "1.0.2" axum-plugin = { git = "https://git.fasharp.io/fa-sharp/axum-plugin", - rev = "9f72278b3c" + rev = "be17dc9aec", + features = ["figment"] } chrono = { version = "0.4.45", @@ -43,7 +44,6 @@ diesel-derive-enum = { version = "3.0.0-beta.1", features = ["postgres"] } diesel-jsonb-derive = { path = "crates/diesel-jsonb-derive" } diesel_migrations = { version = "2.3.2", features = ["postgres"] } dotenvy = "0.15.7" -figment = { version = "0.10.19", features = ["env", "toml"] } fred = { version = "10.1.0", default-features = false, diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index 43cfa83..47ac5cf 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -9,7 +9,7 @@ use axum::{Extension, routing::get}; use axum_plugin::AdHocPlugin; use strum::{Display, EnumIter, EnumMessage, IntoEnumIterator, IntoStaticStr}; -use crate::state::AppState; +use crate::{config::AppConfig, state::AppState}; pub mod api_key; pub mod auth; @@ -34,8 +34,8 @@ enum ApiTag { } /// Adds all API routes with OpenAPI docs to the server under `/api/v1` -pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("API routes").on_setup(|router, _state| { +pub fn plugin() -> AdHocPlugin { + AdHocPlugin::named("API routes").on_setup(|_app, router| { let mut openapi = OpenApi::default(); let api_routes = ApiRouter::new() .nest("/api_key", api_key::routes()) @@ -46,9 +46,7 @@ pub fn plugin() -> AdHocPlugin { .nest("/chat", chat::routes()) .nest("/health", health::routes()) .nest("/provider", provider::routes()) - .finish_api_with(&mut openapi, build_openapi_doc); - - let api_routes_with_docs = api_routes + .finish_api_with(&mut openapi, build_openapi_doc) .route( "/docs/openapi.json", get(async |Extension(openapi): Extension>| axum::Json(openapi)) @@ -61,7 +59,7 @@ pub fn plugin() -> AdHocPlugin { .axum_handler()), ); - Ok(router.nest(API_BASE, api_routes_with_docs)) + Ok(router.nest(API_BASE, api_routes)) }) } diff --git a/server-new/src/config.rs b/server-new/src/config.rs index b043fac..8a61818 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -1,32 +1,22 @@ use std::net::IpAddr; -use anyhow::Context; -use axum_plugin::AdHocPlugin; -use figment::providers::{Env, Format, Serialized, Toml}; +use axum_plugin::figment::{ + Figment, + providers::{Env, Format, Serialized, Toml}, +}; use serde::{Deserialize, Serialize}; -use crate::{ - services::auth::{ - oauth::{DiscordOAuthConfig, GitHubOAuthConfig, GoogleOAuthConfig, OidcConfig}, - proxy::ProxyHeaderConfig, - }, - state::AppState, +use crate::services::auth::{ + oauth::{DiscordOAuthConfig, GitHubOAuthConfig, GoogleOAuthConfig, OidcConfig}, + proxy::ProxyHeaderConfig, }; -/// Plugin that reads and validates configuration, and adds it to server state -pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("Config").on_init(async |mut state| { - let config = extract_config()?; - tracing::info!( - log_level = config.server.log_level, - base_url = config.server.base_url, - host = %config.server.host, - port = config.server.port, - "Config loaded!" - ); - state.insert(config); - Ok(state) - }) +/// Extract configuration from defaults, local `config.toml`, then `RS_CHAT_` environment variables split by `__`. +/// See https://docs.rs/figment/latest/figment/index.html#for-application-authors +pub fn figment() -> Figment { + Figment::from(Serialized::defaults(AppConfig::default())) + .merge(Toml::file("config.toml")) + .merge(Env::prefixed("RS_CHAT_").split("__")) } /// Parsed app configuration @@ -143,15 +133,3 @@ impl Default for RedisConfig { } } } - -/// Extract configuration from defaults, local `config.toml`, then `RS_CHAT_` environment variables. -/// See https://docs.rs/figment/latest/figment/index.html#for-application-authors -fn extract_config() -> anyhow::Result { - let config = figment::Figment::from(Serialized::defaults(AppConfig::default())) - .merge(Toml::file("config.toml")) - .merge(Env::prefixed("RS_CHAT_").split("__")) - .extract::() - .context("Failed to extract valid configuration")?; - - Ok(config) -} diff --git a/server-new/src/lib.rs b/server-new/src/lib.rs index 0df07ba..138718a 100644 --- a/server-new/src/lib.rs +++ b/server-new/src/lib.rs @@ -1,6 +1,6 @@ use axum_plugin::{App, InitializedApp}; -use crate::state::AppState; +use crate::{config::AppConfig, plugins::AxumPlugin, state::AppState}; mod api; mod config; @@ -12,9 +12,19 @@ mod plugins; mod services; mod state; -pub async fn create_app() -> anyhow::Result> { - let app = App::new() - .register(config::plugin()) // Extract configuration and add to state +pub async fn create_app() -> anyhow::Result> { + let app = App::from_figment(config::figment())? + .register(AxumPlugin::named("Config").on_init(async |app| { + let config = app.config(); + tracing::info!( + log_level = config.server.log_level, + base_url = config.server.base_url, + host = %config.server.host, + port = config.server.port, + "Config loaded!" + ); + Ok(app) + })) .register(plugins::clients::plugin()) // Initialize HTTP clients .register(plugins::database::plugin()) // Initialize database .register(plugins::redis::plugin()) // Initialize Redis diff --git a/server-new/src/plugins/auth.rs b/server-new/src/plugins/auth.rs index 999826b..80e4c76 100644 --- a/server-new/src/plugins/auth.rs +++ b/server-new/src/plugins/auth.rs @@ -1,47 +1,42 @@ use std::time::Duration; use anyhow::{Context, bail}; -use axum_plugin::AdHocPlugin; use tower_sessions::{CachingSessionStore, Expiry, SessionManagerLayer, cookie}; use tower_sessions_redis_store::RedisStore; use crate::{ - config::AppConfig, db::DbPool, + plugins::AxumPlugin, services::auth::{ encryption::Encryptor, oauth::OAuthService, session::AuthSessionService, session_store::SessionDbStore, }, - state::AppState, }; const REDIS_PREFIX: &str = "rs-chat:sess:"; const CLEANUP_INTERVAL: Duration = Duration::from_mins(15); /// Add auth & session handling to the server. Sessions are stored in Postgres and cached in Redis. -pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("Auth") - .on_init(async |mut state| { - let config = state.get::().context("no config")?; - let http_client = state.get::().context("no HTTP client")?; - let db_pool = state.get::().context("no db pool")?.to_owned(); - +pub fn plugin() -> AxumPlugin { + AxumPlugin::named("Auth") + .on_init(async |mut app| { // Verify encryption key and build encryptor - let encryption_key = hex::decode(&config.auth.encryption_key) + let encryption_key = hex::decode(&app.config().auth.encryption_key) .context("encryption_key must be hex value")?; if encryption_key.len() != 32 { bail!("encryption_key must be 32 bytes"); } let encryptor = Encryptor::new(&encryption_key)?; + app.insert(encryptor)?; // Build configured OAuth providers - let oauth_providers = OAuthService::build_provider_map(config, http_client) + let http_client = app.get::().context("no HTTP client")?; + let oauth_providers = OAuthService::build_provider_map(app.config(), http_client) .context("build OAuth providers")?; - - state.insert(encryptor); - state.insert(oauth_providers); + app.insert(oauth_providers)?; // Start session cleanup task + let db_pool = app.get::().context("no db pool")?.to_owned(); tokio::spawn(async move { let mut interval = tokio::time::interval(CLEANUP_INTERVAL); interval.tick().await; @@ -55,20 +50,21 @@ pub fn plugin() -> AdHocPlugin { } }); - Ok(state) + Ok(app) }) - .on_setup(|router, state: &AppState| { + .on_setup(|app, router| { // Session persistence - let redis_store = RedisStore::with_prefix(state.redis.clone(), REDIS_PREFIX.to_owned()); - let db_store = SessionDbStore::new(state.db_pool.clone()); + let redis_store = + RedisStore::with_prefix(app.state().redis.clone(), REDIS_PREFIX.to_owned()); + let db_store = SessionDbStore::new(app.state().db_pool.clone()); let session_store = CachingSessionStore::new(redis_store, db_store); // Add session / cookie management to router let session_layer = SessionManagerLayer::new(session_store) - .with_name(state.config.auth.cookie_name.clone()) + .with_name(app.config().auth.cookie_name.clone()) .with_expiry(Expiry::OnInactivity(cookie::time::Duration::minutes(15))) // default short session for login/OAuth .with_private(cookie::Key::derive_from(&hex::decode( - &state.config.auth.encryption_key, + &app.config().auth.encryption_key, )?)) .with_path("/") .with_secure(true) diff --git a/server-new/src/plugins/clients.rs b/server-new/src/plugins/clients.rs index a3bb2d3..6e441c0 100644 --- a/server-new/src/plugins/clients.rs +++ b/server-new/src/plugins/clients.rs @@ -1,24 +1,22 @@ use std::time::Duration; use anyhow::Context; -use axum_plugin::AdHocPlugin; -use crate::{config::AppConfig, services::stream::tinistream::TinistreamClient, state::AppState}; +use crate::{plugins::AxumPlugin, services::stream::tinistream::TinistreamClient}; // Default timeout for HTTP requests const TIMEOUT: Duration = Duration::from_secs(10); /// Setup HTTP clients for interacting with LLMs and services -pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("Clients").on_init(async |mut state| { - let config = state.get::().context("no config")?; - +pub fn plugin() -> AxumPlugin { + AxumPlugin::named("Clients").on_init(async |mut app| { // Main HTTP client for LLM provider and OAuth requests. // No total request timeout to allow for long-lived streaming responses. let http_client = reqwest::ClientBuilder::new() .connect_timeout(TIMEOUT) .redirect(reqwest::redirect::Policy::none()) .build()?; + app.insert(http_client)?; // `tinistream` client with API key header let tini_http_client = reqwest::ClientBuilder::new() @@ -26,17 +24,17 @@ pub fn plugin() -> AdHocPlugin { .timeout(TIMEOUT) .default_headers({ let mut headers = reqwest::header::HeaderMap::new(); - headers.insert("X-API-KEY", config.services.streamer_api_key.parse()?); + headers.insert("X-API-KEY", app.config().services.streamer_api_key.parse()?); headers }) .build()?; let tinistream = TinistreamClient::new(tinistream_client::Client::new_with_client( - &config.services.streamer_url, + &app.config().services.streamer_url, tini_http_client, )); + tinistream.ping().await.context("connect to tinistream")?; + app.insert(tinistream)?; - state.insert(http_client); - state.insert(tinistream); - Ok(state) + Ok(app) }) } diff --git a/server-new/src/plugins/database.rs b/server-new/src/plugins/database.rs index e98ef47..bd92ab8 100644 --- a/server-new/src/plugins/database.rs +++ b/server-new/src/plugins/database.rs @@ -1,22 +1,19 @@ use anyhow::Context; -use axum_plugin::AdHocPlugin; use diesel_async::{ AsyncMigrationHarness, AsyncPgConnection, pooled_connection::{AsyncDieselConnectionManager, ManagerConfig, deadpool::Pool}, }; use diesel_migrations::{EmbeddedMigrations, MigrationHarness}; -use crate::{config::AppConfig, db::DbPool, state::AppState}; +use crate::{db::DbPool, plugins::AxumPlugin}; const MIGRATIONS: EmbeddedMigrations = diesel_migrations::embed_migrations!(); -pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("Database") - .on_init(async |mut state| { - let app_config = state.get::().context("missing config")?; - +pub fn plugin() -> AxumPlugin { + AxumPlugin::named("Database") + .on_init(async |mut app| { let manager = AsyncDieselConnectionManager::::new_with_config( - &app_config.database.url, + &app.config().database.url, { let mut config = ManagerConfig::default(); config.recycling_method = @@ -40,15 +37,12 @@ pub fn plugin() -> AdHocPlugin { Err(err) => anyhow::bail!(format!("Migrations failed: {err}")), }; - state.insert(pool); - Ok(state) + app.insert(pool)?; + Ok(app) }) - .on_shutdown(|state: &AppState| { - let pool = state.db_pool.clone(); - async move { - pool.close(); - tracing::info!("Shut down database pool"); - Ok(()) - } + .on_shutdown(async |app| { + app.state().db_pool.close(); + tracing::info!("Shut down database pool"); + Ok(()) }) } diff --git a/server-new/src/plugins/logging.rs b/server-new/src/plugins/logging.rs index a821105..94b13de 100644 --- a/server-new/src/plugins/logging.rs +++ b/server-new/src/plugins/logging.rs @@ -2,7 +2,6 @@ use std::str::FromStr; use anyhow::Context; use axum::{extract::Request, http::HeaderName}; -use axum_plugin::AdHocPlugin; use tower::ServiceBuilder; use tower_http::{ request_id::{MakeRequestUuid, PropagateRequestIdLayer, SetRequestIdLayer}, @@ -10,12 +9,12 @@ use tower_http::{ }; use tracing::Level; -use crate::state::AppState; +use crate::plugins::AxumPlugin; -pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("Request logs").on_setup(|router, state: &AppState| { +pub fn plugin() -> AxumPlugin { + AxumPlugin::named("Request logs").on_setup(|app, router| { const LOG_LEVEL: Level = Level::INFO; - let request_id_header = HeaderName::from_str(&state.config.server.request_id_header) + let request_id_header = HeaderName::from_str(&app.config().server.request_id_header) .context("invalid request ID header")?; let trace_layer = TraceLayer::new_for_http() @@ -35,7 +34,7 @@ pub fn plugin() -> AdHocPlugin { let logging_service = ServiceBuilder::new() .layer(SetRequestIdLayer::new( request_id_header.clone(), - MakeRequestUuid::default(), + MakeRequestUuid, )) .layer(trace_layer) .layer(PropagateRequestIdLayer::new(request_id_header)); diff --git a/server-new/src/plugins/mod.rs b/server-new/src/plugins/mod.rs index 278b9f2..fe892b2 100644 --- a/server-new/src/plugins/mod.rs +++ b/server-new/src/plugins/mod.rs @@ -1,6 +1,13 @@ +use axum_plugin::AdHocPlugin; + +use crate::{config::AppConfig, state::AppState}; + pub mod auth; pub mod clients; pub mod database; pub mod logging; pub mod redis; pub mod security; + +/// Shared plugin type with correct state and config type parameters +pub type AxumPlugin = AdHocPlugin; diff --git a/server-new/src/plugins/redis.rs b/server-new/src/plugins/redis.rs index 8c15de7..57d2c55 100644 --- a/server-new/src/plugins/redis.rs +++ b/server-new/src/plugins/redis.rs @@ -1,17 +1,15 @@ use std::time::Duration; use anyhow::Context; -use axum_plugin::AdHocPlugin; use fred::prelude::*; -use crate::{config::AppConfig, state::AppState}; +use crate::plugins::AxumPlugin; -pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("Redis") - .on_init(async |mut state| { - let app_config = state.get::().context("no config")?; - let config = Config::from_url(&app_config.redis.url).context("invalid Redis URL")?; - let timeout = Duration::from_secs(app_config.redis.timeout); +pub fn plugin() -> AxumPlugin { + AxumPlugin::named("Redis") + .on_init(async |mut app| { + let config = Config::from_url(&app.config().redis.url).context("parse Redis URL")?; + let timeout = Duration::from_secs(app.config().redis.timeout); let pool = Builder::from_config(config) .with_connection_config(|c| { @@ -22,23 +20,21 @@ pub fn plugin() -> AdHocPlugin { .with_performance_config(|c| { c.default_command_timeout = timeout; }) - .build_pool(app_config.redis.pool_size)?; + .build_pool(app.config().redis.pool_size)?; pool.init().await.context("failed to connect to Redis")?; tracing::info!("Connected to Redis"); + app.insert(pool)?; - state.insert(pool); - Ok(state) + Ok(app) }) - .on_shutdown(|state: &AppState| { - let pool = state.redis.clone(); - async move { - if let Err(e) = pool.quit().await { - tracing::warn!("Error shutting down Redis pool: {e}"); - } else { - tracing::info!("Shut down Redis pool") - } - Ok(()) + .on_shutdown(async |app| { + if let Err(e) = app.state().redis.quit().await { + tracing::warn!("Error shutting down Redis pool: {e}"); + } else { + tracing::info!("Shut down Redis pool") } + + Ok(()) }) } diff --git a/server-new/src/plugins/security.rs b/server-new/src/plugins/security.rs index d48917b..1e58a77 100644 --- a/server-new/src/plugins/security.rs +++ b/server-new/src/plugins/security.rs @@ -1,16 +1,15 @@ use std::time::Duration; use axum::http::StatusCode; -use axum_plugin::AdHocPlugin; use tower::ServiceBuilder; use tower_http::{limit::RequestBodyLimitLayer, timeout::TimeoutLayer}; -use crate::state::AppState; +use crate::plugins::AxumPlugin; /// # Security plugin /// Includes body limiter, request timeout, and security headers. -pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("Security").on_setup(|router, state: &AppState| { +pub fn plugin() -> AxumPlugin { + AxumPlugin::named("Security").on_setup(|app, router| { let security_headers = axum_helmet::Helmet::new() .add(axum_helmet::CrossOriginOpenerPolicy::same_origin()) .add(axum_helmet::CrossOriginResourcePolicy::same_origin()) @@ -20,10 +19,10 @@ pub fn plugin() -> AdHocPlugin { .into_layer()?; let service = ServiceBuilder::new() - .layer(RequestBodyLimitLayer::new(state.config.security.body_limit)) + .layer(RequestBodyLimitLayer::new(app.config().security.body_limit)) .layer(TimeoutLayer::with_status_code( StatusCode::REQUEST_TIMEOUT, - Duration::from_secs(state.config.security.request_timeout), + Duration::from_secs(app.config().security.request_timeout), )) .layer(security_headers); diff --git a/server-new/src/services/stream/tinistream.rs b/server-new/src/services/stream/tinistream.rs index 16abaef..0d1be4b 100644 --- a/server-new/src/services/stream/tinistream.rs +++ b/server-new/src/services/stream/tinistream.rs @@ -24,6 +24,17 @@ impl TinistreamClient { Self { client } } + /// Test the connection to the tinistream server + pub async fn ping(&self) -> TiniResult<()> { + match self.client.health().send().await { + Ok(_) => Ok(()), + Err(err) => Err(TiniError { + status: err.status().map(|s| s.as_u16()).unwrap_or(500), + message: err.to_string(), + }), + } + } + /// Returns a list of keys with the given prefix that have an active stream. pub async fn active_streams(&self, prefix: &str) -> TiniResult> { let streams = self diff --git a/server-new/src/state.rs b/server-new/src/state.rs index 41523fd..c58708a 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -22,7 +22,7 @@ pub struct AppState(Arc); #[derive(AppState)] pub struct AppStateInner { - pub config: AppConfig, + pub config: Arc, pub db_pool: DbPool, pub encryptor: Encryptor, pub http_client: reqwest::Client, From aad03033eaa19a6afa1455d8e9fff160254a4920 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 14 Jul 2026 03:34:03 -0400 Subject: [PATCH 079/111] update all deps --- server-new/Cargo.lock | 337 ++++++++++++++++++++++++++---------------- server-new/Cargo.toml | 4 +- 2 files changed, 209 insertions(+), 132 deletions(-) diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index c791db6..c5ed6a3 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -130,9 +130,9 @@ checksum = "2a4385e2e34eb35d6b3efe798b9eb88096925d87726c0798709bf56d9ed84af3" [[package]] name = "arc-swap" -version = "1.9.1" +version = "1.9.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6a3a1fd6f75306b68087b831f025c712524bcb19aad54e557b1129cfa0a2b207" +checksum = "c049c0be4daef0b145cb3555416b3b8ef5b7888a38aea1a3a155801fe7b0810b" dependencies = [ "rustversion", ] @@ -209,9 +209,9 @@ checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" [[package]] name = "aws-lc-rs" -version = "1.17.0" +version = "1.17.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5ec2f1fc3ec205783a5da9a7e6c1509cc69dedf09a1949e412c1e18469326d00" +checksum = "4342d8937fc7e5dd9b1c60292261c0670c882a2cd1719cfc11b1af41731e32ad" dependencies = [ "aws-lc-sys", "zeroize", @@ -219,14 +219,15 @@ dependencies = [ [[package]] name = "aws-lc-sys" -version = "0.41.0" +version = "0.42.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1a2f9779ce85b93ab6170dd940ad0169b5766ff848247aff13bb788b832fe3f4" +checksum = "6d9ceb1da931507a12f4fccea479dccd00da1943e1b4ae72d8e502d707361444" dependencies = [ "cc", "cmake", "dunce", "fs_extra", + "pkg-config", ] [[package]] @@ -391,9 +392,9 @@ checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649" [[package]] name = "bytemuck" -version = "1.25.0" +version = "1.25.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c8efb64bd706a16a1bdde310ae86b351e4d21550d98d056f22f8a7f7a2183fec" +checksum = "d6aedf8ae72766347502cf3cb4f41cf5e9cc37d28bee90f1fdaaae15f9cf9424" [[package]] name = "byteorder" @@ -403,9 +404,9 @@ checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b" [[package]] name = "bytes" -version = "1.12.0" +version = "1.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8ae3f5d315924270530207e2a68396c3cc547f6dca3fbdca317cfb1a51edb593" +checksum = "fc652a48c352aef3ea3aed32080501cf3ef6ed5da78602a020c991775b0aff04" [[package]] name = "bytes-utils" @@ -419,9 +420,9 @@ dependencies = [ [[package]] name = "cc" -version = "1.2.65" +version = "1.2.67" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e228eec9be7c17ccb640b59b36a5cd805ea2a564a4c5e162c2f659fea30d3b96" +checksum = "e17dd265a7d0f31ef544e1b20e03add05d3b45b491b633b10d67145d2acc1a38" dependencies = [ "find-msvc-tools", "jobserver", @@ -441,6 +442,17 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" +[[package]] +name = "chacha20" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "rand_core 0.10.1", +] + [[package]] name = "chrono" version = "0.4.45" @@ -501,6 +513,12 @@ dependencies = [ "memchr", ] +[[package]] +name = "const-oid" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6ef517f0926dd24a1582492c791b6a4818a4d94e789a334894aa15b0d12f55c" + [[package]] name = "cookie" version = "0.18.1" @@ -510,10 +528,10 @@ dependencies = [ "aes-gcm 0.10.3", "base64", "hkdf", - "hmac", + "hmac 0.12.1", "percent-encoding", - "rand 0.8.6", - "sha2", + "rand 0.8.7", + "sha2 0.10.9", "subtle", "time", "version_check", @@ -573,18 +591,18 @@ checksum = "338089f42c427b86394a5ee60ff321da23a5c89c9d89514c829687b26359fcff" [[package]] name = "crossbeam-channel" -version = "0.5.15" +version = "0.5.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "82b8f8f868b36967f9606790d1903570de9ceaf870a7bf9fbbd3016d636a2cb2" +checksum = "d85363c37faeca707aef026efa9f3b34d077bce547e48f770770625c6013679e" dependencies = [ "crossbeam-utils", ] [[package]] name = "crossbeam-utils" -version = "0.8.21" +version = "0.8.22" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" +checksum = "61803da095bee82a81bb1a452ecc25d3b2f1416d1897eb86430c6159ef717c17" [[package]] name = "crypto-common" @@ -738,9 +756,9 @@ dependencies = [ [[package]] name = "diesel" -version = "2.3.10" +version = "2.3.11" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "29fe29a87fb84c631ffb3ba21798c4b1f3a964701ba78f0dce4bf8668562ec88" +checksum = "e54d1f576cd3a3460f212a4615fd12ce1b6303c095b79a44449ffbe627753dc1" dependencies = [ "bitflags", "byteorder", @@ -833,6 +851,18 @@ dependencies = [ "subtle", ] +[[package]] +name = "digest" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1dd6dbb5841937940781866fa1281a1ff7bd3bf827091440879f9994983d5c2" +dependencies = [ + "block-buffer 0.12.1", + "const-oid", + "crypto-common 0.2.2", + "ctutils", +] + [[package]] name = "displaydoc" version = "0.2.6" @@ -969,7 +999,7 @@ dependencies = [ "futures", "log", "parking_lot", - "rand 0.8.6", + "rand 0.8.7", "redis-protocol", "semver", "socket2 0.5.10", @@ -1115,11 +1145,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" dependencies = [ "cfg-if", - "js-sys", "libc", "r-efi 5.3.0", "wasip2", - "wasm-bindgen", ] [[package]] @@ -1129,9 +1157,11 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "300e883d756b2e4ec94e02791f39b04b522276138852cfc41d9fb7e904106099" dependencies = [ "cfg-if", + "js-sys", "libc", "r-efi 6.0.0", "rand_core 0.10.1", + "wasm-bindgen", ] [[package]] @@ -1150,7 +1180,7 @@ version = "0.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2eecf2d5dc9b66b732b97707a0210906b1d30523eb773193ab777c0c84b3e8d5" dependencies = [ - "polyval 0.7.1", + "polyval 0.7.2", ] [[package]] @@ -1195,7 +1225,7 @@ version = "0.12.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7b5f8eb2ad728638ea2c7d47a21db23b7b58a72ed6a38256b8a1849f15fbbdf7" dependencies = [ - "hmac", + "hmac 0.12.1", ] [[package]] @@ -1204,7 +1234,16 @@ version = "0.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" dependencies = [ - "digest", + "digest 0.10.7", +] + +[[package]] +name = "hmac" +version = "0.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6303bc9732ae41b04cb554b844a762b4115a61bfaa81e3e83050991eeb56863f" +dependencies = [ + "digest 0.11.3", ] [[package]] @@ -1219,9 +1258,9 @@ dependencies = [ [[package]] name = "http-body" -version = "1.0.1" +version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1efedce1fb8e6913f23e0c92de8e62cd5b772a67e7b3946df930a62566c93184" +checksum = "ca2a8f2913ee65f60facd6a5905613afaa448497a0230cc41ce022d93290bc2c" dependencies = [ "bytes", "http", @@ -1229,9 +1268,9 @@ dependencies = [ [[package]] name = "http-body-util" -version = "0.1.3" +version = "0.1.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b021d93e26becf5dc7e1b75b1bed1fd93124b374ceb73f43d4d4eafec896a64a" +checksum = "e9f41fd6a08e4d4ec69df65976da761afd5ad5e58a9d4acb46bd1c953a9e3ff2" dependencies = [ "bytes", "futures-core", @@ -1314,7 +1353,7 @@ dependencies = [ "libc", "percent-encoding", "pin-project-lite", - "socket2 0.6.4", + "socket2 0.6.5", "tokio", "tower-service", "tracing", @@ -1552,19 +1591,19 @@ dependencies = [ [[package]] name = "jobserver" -version = "0.1.34" +version = "0.1.35" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9afb3de4395d6b3e67a780b6de64b51c978ecf11cb9a462c66be7d4ca9039d33" +checksum = "1c00acbd29eabad4a2392fa0e921c874934dbbf4194312ad20f04a0ed67a3cb3" dependencies = [ - "getrandom 0.3.4", + "getrandom 0.4.3", "libc", ] [[package]] name = "js-sys" -version = "0.3.102" +version = "0.3.103" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "03d04c30968dffe80775bd4d7fb676131cd04a1fb46d2686dbffbaec2d9dfd31" +checksum = "53b44bfcdb3f8d5837a46dae1ca9660a837176eee74a28b229bc626816589102" dependencies = [ "cfg-if", "futures-util", @@ -1585,9 +1624,9 @@ checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" [[package]] name = "libredox" -version = "0.1.16" +version = "0.1.18" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e02f3bb43d335493c96bf3fd3a321600bf6bd07ed34bc64118e9293bdffea46c" +checksum = "c943259e342f1e06ff2da7a83eabdfe7f92ce10262688dbf1895ff0b3e6e4652" dependencies = [ "libc", ] @@ -1636,19 +1675,19 @@ checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" [[package]] name = "md-5" -version = "0.10.6" +version = "0.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d89e7ee0cfbedfc4da3340218492196241d89eefb6dab27de5df917a6d2e78cf" +checksum = "69b6441f590336821bb897fb28fc622898ccceb1d6cea3fde5ea86b090c4de98" dependencies = [ "cfg-if", - "digest", + "digest 0.11.3", ] [[package]] name = "memchr" -version = "2.8.2" +version = "2.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "88904434abc2901f197fe8cc55f0445e7ded921dba5911dad2e2b39b48e663c4" +checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" [[package]] name = "migrations_internals" @@ -1685,9 +1724,9 @@ checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" [[package]] name = "mio" -version = "1.2.1" +version = "1.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "02bd0af71c67b473010cbbc60715ee815645a4dc942899111f494b4b737d6fda" +checksum = "30d65c71f1ce40ab09135ce117d742b9f8a19ff91a41a8b57ed50bc2de59c427" dependencies = [ "libc", "wasi 0.11.1+wasi-snapshot-preview1", @@ -1748,11 +1787,11 @@ dependencies = [ "chrono", "getrandom 0.2.17", "http", - "rand 0.8.6", + "rand 0.8.7", "serde", "serde_json", "serde_path_to_error", - "sha2", + "sha2 0.10.9", "thiserror 1.0.69", "url", ] @@ -1880,6 +1919,12 @@ version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" +[[package]] +name = "pkg-config" +version = "0.3.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" + [[package]] name = "polyval" version = "0.6.2" @@ -1894,9 +1939,9 @@ dependencies = [ [[package]] name = "polyval" -version = "0.7.1" +version = "0.7.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7dfc63250416fea14f5749b90725916a6c903f599d51cb635aa7a52bfd03eede" +checksum = "b20f20e954175de5f463f67781b35583397d916b1d148738923711b2ad16bee8" dependencies = [ "cpubits", "cpufeatures 0.3.0", @@ -1905,27 +1950,27 @@ dependencies = [ [[package]] name = "postgres-protocol" -version = "0.6.10" +version = "0.6.12" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3ee9dd5fe15055d2b6806f4736aa0c9637217074e224bbec46d4041b91bb9491" +checksum = "08808e3c483c46e999108051c78334f473d5adb59d78bb80a1268c7e6aa6c514" dependencies = [ "base64", "byteorder", "bytes", "fallible-iterator", - "hmac", + "hmac 0.13.0", "md-5", "memchr", - "rand 0.9.4", - "sha2", + "rand 0.10.2", + "sha2 0.11.0", "stringprep", ] [[package]] name = "postgres-types" -version = "0.2.12" +version = "0.2.14" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "54b858f82211e84682fecd373f68e1ceae642d8d751a1ebd13f33de6257b3e20" +checksum = "851ca9db4932932d69f3ea811b1abe63087a0f740a47692619dd40d4899b68be" dependencies = [ "bytes", "fallible-iterator", @@ -2016,7 +2061,7 @@ dependencies = [ "quinn-udp", "rustc-hash", "rustls", - "socket2 0.6.4", + "socket2 0.6.5", "thiserror 2.0.18", "tokio", "tracing", @@ -2025,15 +2070,16 @@ dependencies = [ [[package]] name = "quinn-proto" -version = "0.11.15" +version = "0.11.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4fcb935c5bec503c2f0e306bdd3e58bb9029dcb14fa8d9ac76e3a5256ac0763e" +checksum = "2f4bfc015262b9df63c8845072ce59068853ff5872180c2ce2f13038b970e560" dependencies = [ "aws-lc-rs", "bytes", - "getrandom 0.3.4", + "getrandom 0.4.3", "lru-slab", - "rand 0.9.4", + "rand 0.10.2", + "rand_pcg", "ring", "rustc-hash", "rustls", @@ -2047,16 +2093,16 @@ dependencies = [ [[package]] name = "quinn-udp" -version = "0.5.14" +version = "0.5.15" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "addec6a0dcad8a8d96a771f815f0eaf55f9d1805756410b39f5fa81332574cbd" +checksum = "35a133f956daabe89a61a685c2649f13d82d5aa4bd5d12d1277e1072a21c0694" dependencies = [ "cfg_aliases", "libc", "once_cell", - "socket2 0.6.4", + "socket2 0.6.5", "tracing", - "windows-sys 0.52.0", + "windows-sys 0.61.2", ] [[package]] @@ -2082,9 +2128,9 @@ checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" [[package]] name = "rand" -version = "0.8.6" +version = "0.8.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5ca0ecfa931c29007047d1bc58e623ab12e5590e8c7cc53200d5202b69266d8a" +checksum = "22f6172bdec972074665ed81ed53b71da00bfc44b65a753cfde883ec4c702a1a" dependencies = [ "libc", "rand_chacha 0.3.1", @@ -2093,14 +2139,25 @@ dependencies = [ [[package]] name = "rand" -version = "0.9.4" +version = "0.9.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "44c5af06bb1b7d3216d91932aed5265164bf384dc89cd6ba05cf59a35f5f76ea" +checksum = "b9ef1d0d795eb7d84685bca4f72f3649f064e6641543d3a8c415898726a57b41" dependencies = [ "rand_chacha 0.9.0", "rand_core 0.9.5", ] +[[package]] +name = "rand" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80" +dependencies = [ + "chacha20", + "getrandom 0.4.3", + "rand_core 0.10.1", +] + [[package]] name = "rand_chacha" version = "0.3.1" @@ -2145,6 +2202,15 @@ version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" +[[package]] +name = "rand_pcg" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "caa0f4137e1c0a72f4c651489402276c8e8e1cf081f3b0ba156d2cbeef09e86a" +dependencies = [ + "rand_core 0.10.1", +] + [[package]] name = "redis-protocol" version = "6.0.0" @@ -2190,9 +2256,9 @@ dependencies = [ [[package]] name = "regex-automata" -version = "0.4.14" +version = "0.4.15" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6e1dd4122fc1595e8162618945476892eefca7b88c52820e74af6262213cae8f" +checksum = "1f388202e4b80542a0921078cc23b6333bcf1409c1e3f86404cae4766a6131db" dependencies = [ "aho-corasick", "memchr", @@ -2347,9 +2413,9 @@ dependencies = [ [[package]] name = "rustc-hash" -version = "2.1.2" +version = "2.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "94300abf3f1ae2e2b8ffb7b58043de3d399c73fa6f4b73826402a5c457614dbe" +checksum = "6b1e7f9a428571be2dc5bc0505c13fb6bf936822b894ec87abf8a08a4e51742d" [[package]] name = "rustc_version" @@ -2362,9 +2428,9 @@ dependencies = [ [[package]] name = "rustls" -version = "0.23.41" +version = "0.23.42" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6b92b125634d9b795e7beca796cc790df15a7fb38323bf3196fda83292d06b1f" +checksum = "3c54fcab019b409d04215d3a17cb438fd7fbf192ee61461f20f4fe18704bc138" dependencies = [ "aws-lc-rs", "once_cell", @@ -2388,9 +2454,9 @@ dependencies = [ [[package]] name = "rustls-pki-types" -version = "1.14.1" +version = "1.15.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "30a7197ae7eb376e574fe940d068c30fe0462554a3ddbe4eca7838e049c937a9" +checksum = "764899a24af3980067ee14bc143654f297b22eaebfe3c7b6b211920a5a59b046" dependencies = [ "web-time", "zeroize", @@ -2437,9 +2503,9 @@ dependencies = [ [[package]] name = "rustversion" -version = "1.0.22" +version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d" +checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f" [[package]] name = "ryu" @@ -2661,13 +2727,13 @@ dependencies = [ [[package]] name = "sha1" -version = "0.10.6" +version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" +checksum = "a978451301f4db1d02937a4ab3ccce137717b81826e79b7d49ffe3244a13c3b8" dependencies = [ "cfg-if", "cpufeatures 0.2.17", - "digest", + "digest 0.10.7", ] [[package]] @@ -2678,7 +2744,18 @@ checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" dependencies = [ "cfg-if", "cpufeatures 0.2.17", - "digest", + "digest 0.10.7", +] + +[[package]] +name = "sha2" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "446ba717509524cb3f22f17ecc096f10f4822d76ab5c0b9822c5f9c284e825f4" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "digest 0.11.3", ] [[package]] @@ -2708,9 +2785,9 @@ dependencies = [ [[package]] name = "simd_cesu8" -version = "1.1.1" +version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "94f90157bb87cddf702797c5dadfa0be7d266cdf49e22da2fcaa32eff75b2c33" +checksum = "11031e251abf8611c80f460e19dbdeb54a66db918e49c65a7065b46ac7aec520" dependencies = [ "rustc_version", "simdutf8", @@ -2739,9 +2816,9 @@ dependencies = [ [[package]] name = "siphasher" -version = "1.0.2" +version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b2aa850e253778c88a04c3d7323b043aeda9d3e30d5971937c1855769763678e" +checksum = "8ee5873ec9cce0195efcb7a4e9507a04cd49aec9c83d0389df45b1ef7ba2e649" [[package]] name = "slab" @@ -2767,9 +2844,9 @@ dependencies = [ [[package]] name = "socket2" -version = "0.6.4" +version = "0.6.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "52d1cfed4120b4d927bf7c0f86d2087a4a7d6027c906d9f9d525a80573b9be51" +checksum = "c3d1e2c7f27f8d4cb10542a02c49005dbd6e93095799d6f3be745fae9f8fedd4" dependencies = [ "libc", "windows-sys 0.61.2", @@ -2904,18 +2981,18 @@ dependencies = [ [[package]] name = "thread_local" -version = "1.1.9" +version = "1.1.10" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f60246a4944f24f6e018aa17cdeffb7818b76356965d03b07d6a9886e8962185" +checksum = "1ad99c4c6d32803332c548b1af0540b357b3f5fc0be8f6c6bfe8b2e6ae784070" dependencies = [ "cfg-if", ] [[package]] name = "time" -version = "0.3.51" +version = "0.3.53" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "85c17d80feb7334b40c484e45ed1a5273dfd8bfda537c3be2e74a06a6686f327" +checksum = "18dfaaeddcb932337b5e7866ee7d0ce9b76d2fd092997146f187ec09b4558a50" dependencies = [ "deranged", "num-conv", @@ -2933,9 +3010,9 @@ checksum = "9e1c906769ad99c88eaa54e728060edef082f8e358ff32030cb7c7d315e81109" [[package]] name = "time-macros" -version = "0.2.30" +version = "0.2.31" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dcef1a61bdb119096e153208ec5cbec23944ce8bca13be5c7f60c634f7403935" +checksum = "c431b87111666e491a90baa837f914fb45cd5dc3c268591b0220ff5057f2085f" dependencies = [ "num-conv", "time-core", @@ -2966,9 +3043,9 @@ dependencies = [ [[package]] name = "tinyvec" -version = "1.11.0" +version = "1.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3e61e67053d25a4e82c844e8424039d9745781b3fc4f32b8d55ed50f5f667ef3" +checksum = "bb4ebadaa0af04fab11ae01eb5f9fdb5f9c5b875506e210e71c07873528baa7f" dependencies = [ "tinyvec_macros", ] @@ -2990,7 +3067,7 @@ dependencies = [ "mio", "pin-project-lite", "signal-hook-registry", - "socket2 0.6.4", + "socket2 0.6.5", "tokio-macros", "windows-sys 0.61.2", ] @@ -3008,9 +3085,9 @@ dependencies = [ [[package]] name = "tokio-postgres" -version = "0.7.16" +version = "0.7.18" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dcea47c8f71744367793f16c2db1f11cb859d28f436bdb4ca9193eb1f787ee42" +checksum = "a528f7d280f6d5b9cd149635c8705b0dd049754bc67d81d31fa25169a93809d3" dependencies = [ "async-trait", "byteorder", @@ -3025,8 +3102,8 @@ dependencies = [ "pin-project-lite", "postgres-protocol", "postgres-types", - "rand 0.9.4", - "socket2 0.6.4", + "rand 0.10.2", + "socket2 0.6.5", "tokio", "tokio-util", "whoami", @@ -3130,7 +3207,7 @@ version = "1.1.2+spec-1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a2abe9b86193656635d2411dc43050282ca48aa31c2451210f4202550afb7526" dependencies = [ - "winnow 1.0.3", + "winnow 1.0.4", ] [[package]] @@ -3251,7 +3328,7 @@ dependencies = [ "futures", "http", "parking_lot", - "rand 0.9.4", + "rand 0.9.5", "serde", "serde_json", "thiserror 2.0.18", @@ -3390,7 +3467,7 @@ dependencies = [ "http", "httparse", "log", - "rand 0.9.4", + "rand 0.9.5", "sha1", "thiserror 2.0.18", "utf-8", @@ -3506,9 +3583,9 @@ checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" [[package]] name = "uuid" -version = "1.23.4" +version = "1.23.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bf80a72845275afea99e7f2b434723d3bc7e38470fcd1c7ed39a599c73319a53" +checksum = "ea5fab0d6c3c01ae70085a09cb03d4c7a1d6314e2b3e075392783396d724ca0a" dependencies = [ "getrandom 0.4.3", "js-sys", @@ -3582,9 +3659,9 @@ dependencies = [ [[package]] name = "wasm-bindgen" -version = "0.2.125" +version = "0.2.126" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8ddb3f79143bced6de84270411622a2699cee572fc0875aeaf1e7867cf9fca1a" +checksum = "4b067c0c11094aef6b7a801c1e34a26affafdf3d051dba08456b868789aaf9a4" dependencies = [ "cfg-if", "once_cell", @@ -3595,9 +3672,9 @@ dependencies = [ [[package]] name = "wasm-bindgen-futures" -version = "0.4.75" +version = "0.4.76" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "503b14d284f2c8dac03b819967e155ea753f573586193b2b2c95990cb5d69280" +checksum = "c62df1340f32221cb9c54d6a27b030e3dba64361d4a95bed55f9aacb44da291d" dependencies = [ "js-sys", "wasm-bindgen", @@ -3605,9 +3682,9 @@ dependencies = [ [[package]] name = "wasm-bindgen-macro" -version = "0.2.125" +version = "0.2.126" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4e21a184b13fb19e157296e2c46056aec9092264fab83e4ba59e68c61b323c3d" +checksum = "167ce5e579f6bcf889c4f7175a8a5a585de84e8ff93976ce393efa5f2837aab1" dependencies = [ "quote", "wasm-bindgen-macro-support", @@ -3615,9 +3692,9 @@ dependencies = [ [[package]] name = "wasm-bindgen-macro-support" -version = "0.2.125" +version = "0.2.126" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fecefd9c35bd935a20fc3fc344b5f29138961e4f47fb03297d88f2587afb5ebd" +checksum = "f3997c7839262f4ef12cf90b818d6340c18e80f263f1a94bf157d0ec4420380e" dependencies = [ "bumpalo", "proc-macro2", @@ -3628,9 +3705,9 @@ dependencies = [ [[package]] name = "wasm-bindgen-shared" -version = "0.2.125" +version = "0.2.126" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "23939e44bb9a5d7576fa2b563dc2e136628f1224e88a8deed09e04858b77871f" +checksum = "dc1b4cb0cc549fcf58d7dfc081778139b3d283a081644e833e84682ad71cea24" dependencies = [ "unicode-ident", ] @@ -3650,9 +3727,9 @@ dependencies = [ [[package]] name = "web-sys" -version = "0.3.102" +version = "0.3.103" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a6430a72df5eb332242960fe84b3002a241163998241eb596d4f739b9757061d" +checksum = "8622dcb61c0bcc9fffa6938bed81210af2da9a7e4a1a834b2e37a59b6dfb6141" dependencies = [ "js-sys", "wasm-bindgen", @@ -3679,9 +3756,9 @@ dependencies = [ [[package]] name = "whoami" -version = "2.1.1" +version = "2.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d6a5b12f9df4f978d2cfdb1bd3bac52433f44393342d7ee9c25f5a1c14c0f45d" +checksum = "998767ef88740d1f5b0682a9c53c24431453923962269c2db68ee43788c5a40d" dependencies = [ "libc", "libredox", @@ -3851,9 +3928,9 @@ dependencies = [ [[package]] name = "winnow" -version = "1.0.3" +version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0592e1c9d151f854e6fd382574c3a0855250e1d9b2f99d9281c6e6391af352f1" +checksum = "23b97319f7b8343df12cc98938e5c3eb436064524c8d2b4e30a1d3a36eecdf81" [[package]] name = "wit-bindgen" @@ -3898,18 +3975,18 @@ dependencies = [ [[package]] name = "zerocopy" -version = "0.8.52" +version = "0.8.54" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ce1022995ff5ff5d841ad7d994facc23098cd40152f2c1d11cd607c6f530653f" +checksum = "b7cbbc0a705a0fd05cc3676525980d2bf5a9bc4adac6d6475209a7887cf59d19" dependencies = [ "zerocopy-derive", ] [[package]] name = "zerocopy-derive" -version = "0.8.52" +version = "0.8.54" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1ae7f38b72ec2a254e2b87ef277cf2cd4fb97cbebf944faa6f33354da0867930" +checksum = "e2e817b7b52d0c7358d3246da9d69935ebb18116b2b102b4230dac079b4862f5" dependencies = [ "proc-macro2", "quote", @@ -3978,6 +4055,6 @@ dependencies = [ [[package]] name = "zmij" -version = "1.0.21" +version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" +checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index ba931b5..8224ef2 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -32,7 +32,7 @@ chrono = { features = ["now", "serde", "std"] } diesel = { - version = "2.3.10", + version = "2.3.11", default-features = false, features = ["chrono", "serde_json", "uuid"] } @@ -103,4 +103,4 @@ tower-sessions-redis-store = { tracing = "0.1.44" tracing-appender = "0.2.5" tracing-subscriber = { version = "0.3.23", features = ["env-filter", "json"] } -uuid = { version = "1.23.4", features = ["serde", "v4"] } +uuid = { version = "1.23.5", features = ["serde", "v4"] } From 67a1119c3fcb065f2982710e0f253a743a0376fc Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 14 Jul 2026 16:28:00 -0400 Subject: [PATCH 080/111] readability tweaks --- server-new/src/services/model/mod.rs | 8 ++++---- server-new/src/services/stream/mod.rs | 8 ++++---- server-new/src/services/stream/tinistream.rs | 14 +++++--------- 3 files changed, 13 insertions(+), 17 deletions(-) diff --git a/server-new/src/services/model/mod.rs b/server-new/src/services/model/mod.rs index f28326d..d746045 100644 --- a/server-new/src/services/model/mod.rs +++ b/server-new/src/services/model/mod.rs @@ -13,10 +13,6 @@ pub mod error; mod providers; pub mod types; -const MODELS_DEV_URL: &str = "https://models.dev/api.json"; -const CACHE_KEY: &str = "rs-chat:models"; -const CACHE_TTL: i64 = 86400; // 1 day in seconds - /// Service for fetching/listing available LLM models pub struct ModelService<'r> { redis: &'r fred::prelude::Pool, @@ -61,6 +57,10 @@ impl<'r> ModelService<'r> { &self, md_provider: ModelsDevProvider, ) -> Result, ModelError> { + const MODELS_DEV_URL: &str = "https://models.dev/api.json"; + const CACHE_KEY: &str = "rs-chat:models"; + const CACHE_TTL: i64 = 86400; // 1 day in seconds + if let Some(models) = self .redis .hget::, _, _>(CACHE_KEY, md_provider.as_ref()) diff --git a/server-new/src/services/stream/mod.rs b/server-new/src/services/stream/mod.rs index 4a04a98..fa9d27c 100644 --- a/server-new/src/services/stream/mod.rs +++ b/server-new/src/services/stream/mod.rs @@ -52,12 +52,12 @@ impl<'r> StreamingService<'r> { format!("user:{}:chat:", user_id) } - /// Check for existing client stream in `tinistream` + /// Check for existing client stream pub async fn exists_stream(&self, stream_key: &str) -> Result { Ok(self.tinistream.stream_exists(&stream_key).await?) } - /// Start the client stream in `tinistream`, and return a WebSocket writer and reader for it + /// Start the client stream, and return a WebSocket writer and reader for it pub async fn create_stream( &self, stream_key: &str, @@ -68,7 +68,7 @@ impl<'r> StreamingService<'r> { Ok((stream_access, writer, reader)) } - /// Process and write the LLM stream to `tinistream` via the WebSocket connection, + /// Process and write the LLM stream via the WebSocket connection, /// and return the accumulated response. pub async fn process_stream( stream: LlmStream, @@ -80,7 +80,7 @@ impl<'r> StreamingService<'r> { .await } - /// Signal end of stream in `tinistream` + /// Signal end of stream pub async fn end_stream(&self, stream_key: &str) -> Result { Ok(self.tinistream.stream_end(stream_key).await?) } diff --git a/server-new/src/services/stream/tinistream.rs b/server-new/src/services/stream/tinistream.rs index 0d1be4b..639ec46 100644 --- a/server-new/src/services/stream/tinistream.rs +++ b/server-new/src/services/stream/tinistream.rs @@ -72,6 +72,7 @@ impl TinistreamClient { .body(StreamRequest::builder().key(key)) .send() .await?; + Ok(res.into_inner()) } @@ -83,6 +84,7 @@ impl TinistreamClient { .body(StreamRequest::builder().key(key)) .send() .await?; + Ok(res.into_inner()) } @@ -95,6 +97,7 @@ impl TinistreamClient { .upgrade() .send() .await?; + res.into_websocket().await } @@ -106,6 +109,7 @@ impl TinistreamClient { .body(StreamRequest::builder().key(key)) .send() .await?; + Ok(res.into_inner().status) } @@ -117,6 +121,7 @@ impl TinistreamClient { .body(StreamRequest::builder().key(key)) .send() .await?; + Ok(res.into_inner().status) } } @@ -139,12 +144,3 @@ impl From> for TiniError { } } } - -impl From for TiniError { - fn from(value: error::ConversionError) -> Self { - TiniError { - status: 400, - message: value.to_string(), - } - } -} From 3a985906ee3356684e5f7204d47c601c620782f5 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 14 Jul 2026 19:08:16 -0400 Subject: [PATCH 081/111] add chat session routes --- docker-compose.yml | 32 ++--- server-new/src/api/mod.rs | 2 + server-new/src/api/session.rs | 159 +++++++++++++++++++++++++ server-new/src/db/mod.rs | 2 +- server-new/src/db/models/chat.rs | 16 +-- server-new/src/db/queries.rs | 3 +- server-new/src/db/repositories/chat.rs | 52 ++++---- server-new/src/llm/types.rs | 2 +- server-new/src/services/chat/mod.rs | 2 +- 9 files changed, 214 insertions(+), 56 deletions(-) create mode 100644 server-new/src/api/session.rs diff --git a/docker-compose.yml b/docker-compose.yml index 2974726..04fa64a 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -12,7 +12,7 @@ services: - postgres_data:/var/lib/postgresql/data redis: - image: valkey/valkey:7-alpine + image: valkey/valkey:8-alpine container_name: redis ports: - "6379:6379" @@ -33,21 +33,21 @@ services: depends_on: - redis - runner: - image: ghcr.io/fa-sharp/tinirun:0.1.1 - container_name: tinirun - ports: - - "8082:8082" - environment: - RUNNER_HOST: 0.0.0.0 - RUNNER_PORT: 8082 - RUNNER_LOG_LEVEL: info - RUNNER_REDIS_URL: redis://redis:6379 - env_file: server/.env - depends_on: - - redis - volumes: - - /var/run/docker.sock:/var/run/docker.sock + # runner: + # image: ghcr.io/fa-sharp/tinirun:0.1.1 + # container_name: tinirun + # ports: + # - "8082:8082" + # environment: + # RUNNER_HOST: 0.0.0.0 + # RUNNER_PORT: 8082 + # RUNNER_LOG_LEVEL: info + # RUNNER_REDIS_URL: redis://redis:6379 + # env_file: server/.env + # depends_on: + # - redis + # volumes: + # - /var/run/docker.sock:/var/run/docker.sock # rschat: # build: diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index 47ac5cf..2ac879a 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -16,6 +16,7 @@ pub mod auth; pub mod chat; pub mod health; pub mod provider; +pub mod session; const API_BASE: &str = "/api/v1"; const API_AUTH_BASE: &str = "/api/v1/auth"; @@ -46,6 +47,7 @@ pub fn plugin() -> AdHocPlugin { .nest("/chat", chat::routes()) .nest("/health", health::routes()) .nest("/provider", provider::routes()) + .nest("/session", session::routes()) .finish_api_with(&mut openapi, build_openapi_doc) .route( "/docs/openapi.json", diff --git a/server-new/src/api/session.rs b/server-new/src/api/session.rs new file mode 100644 index 0000000..3eb368e --- /dev/null +++ b/server-new/src/api/session.rs @@ -0,0 +1,159 @@ +use std::borrow::Cow; + +use axum::{ + Json, + extract::{Path, Query}, +}; +use axum_aide_macros::api_routes; +use schemars::JsonSchema; +use serde::{Deserialize, Serialize}; +use uuid::Uuid; + +use crate::{ + api::ApiTag, + db::{ + models::{ChatRsMessage, ChatRsSession, NewChatRsSession, UpdateChatRsSession}, + queries::FullTextSearchResult, + }, + error::{AppError, AppResult}, + extractors::{CurrentUser, Database}, + services::chat::DEFAULT_SESSION_TITLE, + state::AppState, +}; + +api_routes! { + state: AppState, + tag: ApiTag::Chat.into(), + GET "/" => get_recent_sessions, "List recent chat sessions"; + POST "/" => create_session, "Create chat session"; + GET "/{session_id}" => get_session, "Get chat session"; + GET "/search" => search_sessions, "Search chat sessions"; + PATCH "/{session_id}" => update_session, "Update chat session"; + DELETE "/{session_id}" => delete_session, "Delete chat session"; + DELETE "/{session_id}/{message_id}" => delete_message, "Delete chat message"; +} + +async fn get_recent_sessions( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, +) -> AppResult>> { + let sessions = db.chats().get_recent_sessions(&user_id).await?; + + Ok(Json(sessions)) +} + +async fn create_session( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, +) -> AppResult> { + let session_id = db + .chats() + .create_session(NewChatRsSession { + user_id: &user_id, + title: DEFAULT_SESSION_TITLE, + }) + .await?; + + Ok(Json(SessionIdResponse { session_id })) +} + +async fn get_session( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, + Path(session_id): Path, +) -> AppResult> { + let (session, messages) = db + .chats() + .find_session_with_messages(&user_id, &session_id) + .await?; + + match session { + Some(session) => Ok(Json(GetSessionResponse { session, messages })), + None => Err(AppError::not_found("chat session not found")), + } +} + +async fn search_sessions( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, + Query(SessionSearchQuery { query }): Query>, +) -> AppResult>> { + let sessions = db.chats().search_sessions(&user_id, &query).await?; + + Ok(Json(sessions)) +} + +async fn update_session( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, + Path(session_id): Path, + Json(input): Json, +) -> AppResult> { + let updated_id = db + .chats() + .update_session( + &user_id, + &session_id, + UpdateChatRsSession { + title: Some(&input.title), + ..Default::default() + }, + ) + .await?; + + match updated_id { + Some(session_id) => Ok(Json(SessionIdResponse { session_id })), + None => Err(AppError::not_found("chat session not found")), + } +} + +async fn delete_session( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, + Path(session_id): Path, +) -> AppResult> { + match db.chats().delete_session(&user_id, &session_id).await? { + Some(session_id) => Ok(Json(SessionIdResponse { session_id })), + None => Err(AppError::not_found("chat session not found")), + } +} + +async fn delete_message( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, + Path((session_id, message_id)): Path<(Uuid, Uuid)>, +) -> AppResult> { + match db.chats().find_session(&user_id, &session_id).await? { + Some(session) => match db.chats().delete_message(&session.id, &message_id).await? { + Some(message_id) => Ok(Json(MessageIdResponse { message_id })), + None => Err(AppError::not_found("chat message not found")), + }, + None => Err(AppError::not_found("chat session not found")), + } +} + +#[derive(Serialize, JsonSchema)] +struct SessionIdResponse { + session_id: Uuid, +} + +#[derive(Serialize, JsonSchema)] +struct MessageIdResponse { + message_id: Uuid, +} + +#[derive(Serialize, JsonSchema)] +struct GetSessionResponse { + session: ChatRsSession, + messages: Vec, +} + +#[derive(Deserialize, JsonSchema)] +struct SessionSearchQuery<'q> { + query: Cow<'q, str>, +} + +#[derive(Deserialize, JsonSchema)] +struct UpdateSessionInput { + title: String, +} diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs index 0cffd78..d78ec01 100644 --- a/server-new/src/db/mod.rs +++ b/server-new/src/db/mod.rs @@ -8,7 +8,7 @@ use diesel_async::{ }; pub mod models; -mod queries; +pub mod queries; mod repositories; mod schema; diff --git a/server-new/src/db/models/chat.rs b/server-new/src/db/models/chat.rs index 6301c17..842fcf2 100644 --- a/server-new/src/db/models/chat.rs +++ b/server-new/src/db/models/chat.rs @@ -1,6 +1,7 @@ use chrono::{DateTime, Utc}; use diesel::{deserialize::FromSqlRow, expression::AsExpression, prelude::*}; use diesel_jsonb_derive::AsJsonb; +use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use uuid::Uuid; @@ -9,7 +10,7 @@ use crate::{ llm::types::{LlmChatOptions, LlmUsage}, }; -#[derive(Identifiable, Associations, Queryable, Selectable, Serialize)] +#[derive(Identifiable, Associations, Queryable, Selectable, Serialize, JsonSchema)] #[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] #[diesel(table_name = super::schema::chat_sessions)] pub struct ChatRsSession { @@ -22,7 +23,7 @@ pub struct ChatRsSession { pub updated_at: DateTime, } -#[derive(Debug, Default, Serialize, Deserialize, FromSqlRow, AsExpression, AsJsonb)] +#[derive(Debug, Default, Serialize, Deserialize, JsonSchema, FromSqlRow, AsExpression, AsJsonb)] #[diesel(sql_type = diesel::sql_types::Jsonb)] pub struct ChatRsSessionMeta { // /// User configuration of tools for this session @@ -51,7 +52,8 @@ pub struct UpdateChatRsSession<'r> { #[derive(diesel_derive_enum::DbEnum)] #[db_enum(existing_type_path = "crate::db::schema::sql_types::ChatMessageRole")] -#[derive(Debug, PartialEq, Eq, Serialize)] +#[derive(Debug, PartialEq, Eq, Serialize, JsonSchema)] +#[serde(rename_all = "lowercase")] pub enum ChatRsMessageRole { User, Assistant, @@ -59,7 +61,7 @@ pub enum ChatRsMessageRole { Tool, } -#[derive(Identifiable, Queryable, Selectable, Associations, Serialize)] +#[derive(Identifiable, Queryable, Selectable, Associations, Serialize, JsonSchema)] #[diesel(belongs_to(ChatRsSession, foreign_key = session_id))] #[diesel(table_name = super::schema::chat_messages)] pub struct ChatRsMessage { @@ -71,7 +73,7 @@ pub struct ChatRsMessage { pub created_at: DateTime, } -#[derive(Debug, Default, Serialize, Deserialize, AsExpression, FromSqlRow, AsJsonb)] +#[derive(Debug, Default, Serialize, Deserialize, JsonSchema, AsExpression, FromSqlRow, AsJsonb)] #[diesel(sql_type = diesel::sql_types::Jsonb)] pub struct ChatRsMessageMeta { /// User messages: metadata associated with the user message @@ -99,14 +101,14 @@ impl ChatRsMessageMeta { } } -#[derive(Debug, Default, Serialize, Deserialize)] +#[derive(Debug, Default, Serialize, Deserialize, JsonSchema)] pub struct UserMeta { /// The IDs of the files attached to this message #[serde(skip_serializing_if = "Option::is_none")] pub files: Option>, } -#[derive(Debug, Default, Serialize, Deserialize)] +#[derive(Debug, Default, Serialize, Deserialize, JsonSchema)] pub struct AssistantMeta { /// The ID of the LLM provider used to generate this message pub provider_id: i32, diff --git a/server-new/src/db/queries.rs b/server-new/src/db/queries.rs index 3d8ea35..2bdd408 100644 --- a/server-new/src/db/queries.rs +++ b/server-new/src/db/queries.rs @@ -1,12 +1,13 @@ use diesel::{prelude::QueryableByName, sql_query}; use diesel_async::RunQueryDsl; +use schemars::JsonSchema; use serde::Serialize; use uuid::Uuid; use crate::db::DbConnection; /// Session matches for a full-text search query of chat titles and messages -#[derive(Debug, Clone, QueryableByName, Serialize)] +#[derive(Debug, Clone, QueryableByName, Serialize, JsonSchema)] pub struct FullTextSearchResult { #[diesel(sql_type = diesel::sql_types::Uuid)] pub session_id: Uuid, diff --git a/server-new/src/db/repositories/chat.rs b/server-new/src/db/repositories/chat.rs index f24e70b..f61d73f 100644 --- a/server-new/src/db/repositories/chat.rs +++ b/server-new/src/db/repositories/chat.rs @@ -23,13 +23,13 @@ impl<'a> ChatRepository<'a> { pub async fn create_session( &mut self, session: NewChatRsSession<'_>, - ) -> Result { - let id: Uuid = diesel::insert_into(chat_sessions::table) + ) -> Result { + let id = diesel::insert_into(chat_sessions::table) .values(session) .returning(chat_sessions::id) .get_result(self.db) .await?; - Ok(id.to_string()) + Ok(id) } pub async fn save_message( @@ -60,7 +60,7 @@ impl<'a> ChatRepository<'a> { &mut self, user_id: &Uuid, message_id: &Uuid, - ) -> Result { + ) -> Result, diesel::result::Error> { chat_messages::table .inner_join(chat_sessions::table.on(chat_sessions::id.eq(chat_messages::session_id))) .select(ChatRsMessage::as_select()) @@ -68,20 +68,21 @@ impl<'a> ChatRepository<'a> { .filter(chat_messages::id.eq(message_id)) .get_result(self.db) .await + .optional() } pub async fn delete_message( &mut self, session_id: &Uuid, message_id: &Uuid, - ) -> Result { - let id: Uuid = diesel::delete(chat_messages::table) + ) -> Result, diesel::result::Error> { + diesel::delete(chat_messages::table) .filter(chat_messages::session_id.eq(session_id)) .filter(chat_messages::id.eq(message_id)) .returning(chat_messages::id) .get_result(self.db) - .await?; - Ok(id.to_string()) + .await + .optional() } pub async fn get_recent_sessions( @@ -146,9 +147,7 @@ impl<'a> ChatRepository<'a> { user_id: &Uuid, query: &str, ) -> Result, diesel::result::Error> { - let sessions = full_text_query(self.db, user_id, query, 10).await?; - - Ok(sessions) + full_text_query(self.db, user_id, query, 10).await } pub async fn update_session( @@ -156,41 +155,36 @@ impl<'a> ChatRepository<'a> { user_id: &Uuid, session_id: &Uuid, data: UpdateChatRsSession<'_>, - ) -> Result { - let updated_id: Uuid = diesel::update(chat_sessions::table.find(session_id)) + ) -> Result, diesel::result::Error> { + diesel::update(chat_sessions::table.find(session_id)) .set(data) .filter(chat_sessions::user_id.eq(user_id)) .returning(chat_sessions::id) .get_result(self.db) - .await?; - - Ok(updated_id) + .await + .optional() } pub async fn delete_session( &mut self, user_id: &Uuid, session_id: &Uuid, - ) -> Result { - let id: Uuid = diesel::delete(chat_sessions::table.find(session_id)) + ) -> Result, diesel::result::Error> { + diesel::delete(chat_sessions::table.find(session_id)) .filter(chat_sessions::user_id.eq(user_id)) .returning(chat_sessions::id) .get_result(self.db) - .await?; - - Ok(id) + .await + .optional() } - pub async fn delete_by_user( + pub async fn delete_sessions_by_user( &mut self, user_id: &Uuid, - ) -> Result, diesel::result::Error> { - let ids: Vec = diesel::delete(chat_sessions::table) + ) -> Result { + diesel::delete(chat_sessions::table) .filter(chat_sessions::user_id.eq(user_id)) - .returning(chat_sessions::id) - .get_results(self.db) - .await?; - - Ok(ids) + .execute(self.db) + .await } } diff --git a/server-new/src/llm/types.rs b/server-new/src/llm/types.rs index be697d9..8cd2280 100644 --- a/server-new/src/llm/types.rs +++ b/server-new/src/llm/types.rs @@ -59,7 +59,7 @@ pub struct LlmAssistantMessage { } /// Usage stats from the LLM provider -#[derive(Debug, Default, Serialize, Deserialize)] +#[derive(Debug, Default, Serialize, Deserialize, JsonSchema)] pub struct LlmUsage { pub input_tokens: Option, pub output_tokens: Option, diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index e68208b..9405f73 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -25,7 +25,7 @@ mod error; mod messages; mod titles; -const DEFAULT_SESSION_TITLE: &str = "New Chat"; +pub const DEFAULT_SESSION_TITLE: &str = "New Chat"; pub struct ChatService<'r> { db_pool: &'r DbPool, From fdf0a1c00d6e62cb135101b2b0b25c0b05b1c5d7 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 14 Jul 2026 19:57:45 -0400 Subject: [PATCH 082/111] allow regenerating LLM response --- docker-compose.yml | 1 + server-new/src/api/chat.rs | 42 ++++++++- server-new/src/db/models/chat.rs | 2 +- server-new/src/db/repositories/chat.rs | 35 ++++++++ server-new/src/services/chat/error.rs | 3 + server-new/src/services/chat/mod.rs | 120 ++++++++++++++++++------- server-new/src/services/stream/mod.rs | 6 +- 7 files changed, 174 insertions(+), 35 deletions(-) diff --git a/docker-compose.yml b/docker-compose.yml index 04fa64a..1be9823 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -29,6 +29,7 @@ services: STREAMER_API_KEY: dev-streamer-api-key STREAMER_BASE_URL: http://localhost:8081 STREAMER_REDIS_URL: redis://redis:6379 + STREAMER_KEY_PREFIX: "rs-chat:" env_file: server-new/.env depends_on: - redis diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index 9578e4e..0284746 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -24,6 +24,9 @@ api_routes! { POST "/session/{session_id}" => chat_stream, "Chat", { description: "Send a message in a chat session and stream the response" }; + POST "/session/{session_id}/regenerate" => regenerate_response, "Regenerate response", { + description: "Regenerate the latest assistant response in a chat session" + }; } #[derive(Debug, Deserialize, JsonSchema)] @@ -77,8 +80,8 @@ struct ChatInput { } async fn chat_stream( - Path(session_id): Path, CurrentUser { user_id }: CurrentUser, + Path(session_id): Path, Database(mut db): Database, State(state): State, Json(input): Json, @@ -108,6 +111,43 @@ async fn chat_stream( })) } +#[derive(Debug, Deserialize, JsonSchema)] +struct RegenerateInput { + /// The ID of the provider to chat with + provider_id: i32, + /// Configuration for the provider + options: LlmChatOptions, +} + +async fn regenerate_response( + CurrentUser { user_id }: CurrentUser, + Path(session_id): Path, + Database(mut db): Database, + State(state): State, + Json(input): Json, +) -> Result, AppError> { + let llm_provider = state + .provider_service() + .build_llm_provider(&mut db, &user_id, input.provider_id) + .await?; + let stream_access = state + .chat_service() + .regenerate_response( + &mut db, + user_id, + session_id, + input.provider_id, + llm_provider, + input.options, + ) + .await?; + + Ok(Json(StreamAccess { + url: stream_access.sse_url, + token: stream_access.token, + })) +} + /// The URL and Bearer token to access the SSE stream #[derive(Serialize, JsonSchema)] struct StreamAccess { diff --git a/server-new/src/db/models/chat.rs b/server-new/src/db/models/chat.rs index 842fcf2..8224462 100644 --- a/server-new/src/db/models/chat.rs +++ b/server-new/src/db/models/chat.rs @@ -52,7 +52,7 @@ pub struct UpdateChatRsSession<'r> { #[derive(diesel_derive_enum::DbEnum)] #[db_enum(existing_type_path = "crate::db::schema::sql_types::ChatMessageRole")] -#[derive(Debug, PartialEq, Eq, Serialize, JsonSchema)] +#[derive(Debug, strum::EnumIs, Serialize, JsonSchema)] #[serde(rename_all = "lowercase")] pub enum ChatRsMessageRole { User, diff --git a/server-new/src/db/repositories/chat.rs b/server-new/src/db/repositories/chat.rs index f61d73f..e2f6300 100644 --- a/server-new/src/db/repositories/chat.rs +++ b/server-new/src/db/repositories/chat.rs @@ -142,6 +142,41 @@ impl<'a> ChatRepository<'a> { Ok((session.optional()?, messages?)) } + pub async fn find_messages_before( + &mut self, + user_id: &Uuid, + session_id: &Uuid, + message: &ChatRsMessage, + ) -> Result, diesel::result::Error> { + chat_messages::table + .inner_join(chat_sessions::table.on(chat_sessions::id.eq(chat_messages::session_id))) + .select(ChatRsMessage::as_select()) + .filter(chat_sessions::user_id.eq(user_id)) + .filter(chat_messages::session_id.eq(session_id)) + .filter(chat_messages::created_at.lt(message.created_at)) + .order_by(chat_messages::created_at.asc()) + .load(self.db) + .await + } + + pub async fn has_messages_after( + &mut self, + user_id: &Uuid, + session_id: &Uuid, + message: &ChatRsMessage, + ) -> Result { + let count = chat_messages::table + .inner_join(chat_sessions::table.on(chat_sessions::id.eq(chat_messages::session_id))) + .filter(chat_sessions::user_id.eq(user_id)) + .filter(chat_messages::session_id.eq(session_id)) + .filter(chat_messages::created_at.gt(message.created_at)) + .count() + .get_result::(self.db) + .await?; + + Ok(count > 0) + } + pub async fn search_sessions( &mut self, user_id: &Uuid, diff --git a/server-new/src/services/chat/error.rs b/server-new/src/services/chat/error.rs index 48253dc..ad2510c 100644 --- a/server-new/src/services/chat/error.rs +++ b/server-new/src/services/chat/error.rs @@ -7,6 +7,8 @@ pub enum ChatError { Messages, #[error("session not found")] SessionNotFound, + #[error("no assistant response")] + NoAssistantResponse, #[error("already streaming a response")] AlreadyStreaming, #[error(transparent)] @@ -24,6 +26,7 @@ impl From for AppError { match value { ChatError::Messages => Self::bad_request("invalid messages"), ChatError::SessionNotFound => Self::not_found("chat session not found"), + ChatError::NoAssistantResponse => Self::bad_request("no assistant response found"), ChatError::AlreadyStreaming => Self::bad_request("already streaming this chat session"), ChatError::Request(err) => Self::bad_request(err.to_string()), err => Self::internal(err.into()), diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index 9405f73..82ea3a6 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -32,6 +32,15 @@ pub struct ChatService<'r> { tinistream: &'r TinistreamClient, } +/// Chat stream parameters for response generation +struct ChatStreamParams { + user_id: Uuid, + session_id: Uuid, + provider_id: i32, + chat_options: LlmChatOptions, + replace_message_id: Option, +} + impl<'r> ChatService<'r> { pub fn new(db_pool: &'r DbPool, tinistream: &'r TinistreamClient) -> Self { Self { @@ -64,23 +73,14 @@ impl<'r> ChatService<'r> { user_message: Option, chat_options: LlmChatOptions, ) -> Result { - // Get session and message history - let (chat_session, mut session_messages) = db + let (chat_session, mut messages) = db .chats() .find_session_with_messages(&user_id, &session_id) .await?; let chat_session = chat_session.ok_or(ChatError::SessionNotFound)?; - // Check that we're not already streaming a response for this chat session - let stream_key = StreamingService::chat_stream_key(&user_id, &session_id); - let stream_service = StreamingService::new(self.tinistream); - if stream_service.exists_stream(&stream_key).await? { - return Err(ChatError::AlreadyStreaming); - } - - // Save user message, and generate session title if needed if let Some(user_message) = user_message { - if session_messages.is_empty() && chat_session.title == DEFAULT_SESSION_TITLE { + if messages.is_empty() && chat_session.title == DEFAULT_SESSION_TITLE { titles::generate_title( user_id, session_id, @@ -99,31 +99,88 @@ impl<'r> ChatService<'r> { meta: ChatRsMessageMeta::new_user(UserMeta::default()), }) .await?; - session_messages.push(new_message); + messages.push(new_message); } - // Send the request to the LLM provider and get the stream response - let stream = provider + self.start_assistant_stream( + provider, + messages, + ChatStreamParams { + user_id, + session_id, + provider_id, + chat_options, + replace_message_id: None, + }, + ) + .await + } + + pub async fn regenerate_response( + &self, + db: &mut DbService, + user_id: Uuid, + session_id: Uuid, + provider_id: i32, + provider: Arc, + chat_options: LlmChatOptions, + ) -> Result { + let (chat_session, messages) = db + .chats() + .find_session_with_messages(&user_id, &session_id) + .await?; + chat_session.ok_or(ChatError::SessionNotFound)?; + let assistant_message_id = messages + .iter() + .rev() + .find(|m| m.role.is_assistant()) + .ok_or(ChatError::NoAssistantResponse)? + .id; + + self.start_assistant_stream( + provider, + messages, + ChatStreamParams { + user_id, + session_id, + provider_id, + chat_options, + replace_message_id: Some(assistant_message_id), + }, + ) + .await + } + + async fn start_assistant_stream( + &self, + provider: Arc, + messages: Vec, + params: ChatStreamParams, + ) -> Result { + let stream_key = StreamingService::chat_stream_key(¶ms.user_id, ¶ms.session_id); + let stream_service = StreamingService::new(self.tinistream); + if stream_service.exists_stream(&stream_key).await? { + return Err(ChatError::AlreadyStreaming); + } + + let llm_messages = messages::build_llm_messages(messages)?; + let response_stream = provider .stream_chat(LlmChatRequest { - messages: &messages::build_llm_messages(session_messages)?, - options: &chat_options, + messages: &llm_messages, + options: ¶ms.chat_options, }) .await?; - // Create a new client stream in `tinistream` to stream the response to the user let (stream_access, ws_writer, ws_reader) = stream_service.create_stream(&stream_key).await?; - // Spawn task to process and save the streaming response let db_pool = self.db_pool.to_owned(); let tinistream_client = self.tinistream.to_owned(); tokio::spawn(async move { - let response = StreamingService::process_stream(stream, ws_writer, ws_reader).await; + let response = + StreamingService::process_stream(response_stream, ws_writer, ws_reader).await; let stream_cancelled = response.cancelled; - if let Err(err) = - Self::persist_response(db_pool, session_id, provider_id, chat_options, response) - .await - { + if let Err(err) = Self::persist_response(response, params, db_pool).await { tracing::error!("Failed to save assistant response: {err}"); } @@ -140,16 +197,14 @@ impl<'r> ChatService<'r> { /// Save response message and metadata to database async fn persist_response( - db_pool: DbPool, - session_id: Uuid, - provider_id: i32, - chat_options: LlmChatOptions, response: LlmStreamOutput, + params: ChatStreamParams, + db_pool: DbPool, ) -> Result { let mut db = DbService::from_pool(&db_pool).await?; let assistant_meta = AssistantMeta { - provider_id, - provider_options: Some(chat_options), + provider_id: params.provider_id, + provider_options: Some(params.chat_options), // tool_calls: response.tool_calls, // files: image_ids, usage: response.usage, @@ -163,9 +218,14 @@ impl<'r> ChatService<'r> { content: &response.text.unwrap_or_default(), meta: ChatRsMessageMeta::new_assistant(assistant_meta), role: ChatRsMessageRole::Assistant, - session_id: &session_id, + session_id: ¶ms.session_id, }) .await?; + if let Some(message_id) = params.replace_message_id { + db.chats() + .delete_message(¶ms.session_id, &message_id) + .await?; + } Ok(new_message) } diff --git a/server-new/src/services/stream/mod.rs b/server-new/src/services/stream/mod.rs index fa9d27c..eb8cf06 100644 --- a/server-new/src/services/stream/mod.rs +++ b/server-new/src/services/stream/mod.rs @@ -34,8 +34,8 @@ pub struct LlmStreamOutput { pub cancelled: bool, } -pub type WsWriter = SplitSink; -pub type WsReader = SplitStream; +type WsWriter = SplitSink; +type WsReader = SplitStream; impl<'r> StreamingService<'r> { pub fn new(tinistream: &'r TinistreamClient) -> Self { @@ -68,7 +68,7 @@ impl<'r> StreamingService<'r> { Ok((stream_access, writer, reader)) } - /// Process and write the LLM stream via the WebSocket connection, + /// Process and write the LLM response stream via the WebSocket connection, /// and return the accumulated response. pub async fn process_stream( stream: LlmStream, From 666925bb1e90d2970b885e970b44e27523fe0fa5 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 14 Jul 2026 21:05:02 -0400 Subject: [PATCH 083/111] add anthropic and ollama providers --- server-new/src/llm/providers/anthropic/mod.rs | 127 +++++++++++ .../src/llm/providers/anthropic/request.rs | 163 ++++++++++++++ .../src/llm/providers/anthropic/response.rs | 205 ++++++++++++++++++ server-new/src/llm/providers/mod.rs | 4 + server-new/src/llm/providers/ollama/mod.rs | 126 +++++++++++ .../src/llm/providers/ollama/request.rs | 148 +++++++++++++ .../src/llm/providers/ollama/response.rs | 148 +++++++++++++ server-new/src/llm/providers/openai/mod.rs | 21 +- server-new/src/llm/providers/utils.rs | 16 +- server-new/src/services/provider/mod.rs | 26 +-- 10 files changed, 949 insertions(+), 35 deletions(-) create mode 100644 server-new/src/llm/providers/anthropic/mod.rs create mode 100644 server-new/src/llm/providers/anthropic/request.rs create mode 100644 server-new/src/llm/providers/anthropic/response.rs create mode 100644 server-new/src/llm/providers/ollama/mod.rs create mode 100644 server-new/src/llm/providers/ollama/request.rs create mode 100644 server-new/src/llm/providers/ollama/response.rs diff --git a/server-new/src/llm/providers/anthropic/mod.rs b/server-new/src/llm/providers/anthropic/mod.rs new file mode 100644 index 0000000..2a2dc80 --- /dev/null +++ b/server-new/src/llm/providers/anthropic/mod.rs @@ -0,0 +1,127 @@ +//! Anthropic LLM provider + +use futures::StreamExt; + +use crate::llm::{ + error::LlmRequestError, + interface::{LlmPromptResponse, LlmProvider, LlmStreamingResponse}, + providers::utils, + types::{LlmChatRequest, LlmPrompt, LlmUsage}, +}; + +mod request; +mod response; + +use {request::*, response::*}; + +const MESSAGES_API_URL: &str = "https://api.anthropic.com/v1/messages"; +const API_VERSION: &str = "2023-06-01"; +const DEFAULT_MAX_TOKENS: u32 = 4096; + +/// Anthropic chat provider +#[derive(Debug, Clone)] +pub struct AnthropicProvider { + client: reqwest::Client, + api_key: String, +} + +impl AnthropicProvider { + pub fn new(http_client: &reqwest::Client, api_key: impl Into) -> Self { + Self { + client: http_client.clone(), + api_key: api_key.into(), + } + } +} + +impl LlmProvider for AnthropicProvider { + fn prompt<'r>(&'r self, prompt: LlmPrompt<'r>) -> LlmPromptResponse<'r> { + let request = AnthropicRequest { + model: &prompt.options.model, + messages: vec![AnthropicMessage { + role: "user", + content: vec![AnthropicContentBlock::Text { text: prompt.text }], + }], + max_tokens: prompt.options.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS), + temperature: prompt.options.temperature, + system: None, + stream: None, + tools: None, + }; + + Box::pin(async move { + let mut response: AnthropicResponse = utils::llm_api_request( + self.client + .post(MESSAGES_API_URL) + .header("anthropic-version", API_VERSION) + .header("content-type", "application/json") + .header("x-api-key", &self.api_key) + .json(&request), + "Anthropic", + ) + .await? + .json() + .await?; + + let text = response + .content + .get_mut(0) + .and_then(|block| match block { + AnthropicResponseContentBlock::Text { text } => Some(std::mem::take(text)), + _ => None, + }) + .ok_or_else(|| LlmRequestError::NoContent)?; + if let Some(usage) = response.usage { + let usage: LlmUsage = usage.into(); + tracing::info!("Prompt usage: {:?}", usage); + } + + Ok(text) + }) + } + + fn stream_chat<'r>(&'r self, req: LlmChatRequest<'r>) -> LlmStreamingResponse<'r> { + let (anthropic_messages, system_prompt) = build_anthropic_messages(&req.messages); + // let anthropic_tools = tools.as_ref().map(|t| build_anthropic_tools(t)); + let request = AnthropicRequest { + model: &req.options.model, + messages: anthropic_messages, + max_tokens: req.options.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS), + temperature: req.options.temperature, + system: system_prompt, + stream: Some(true), + // tools: anthropic_tools, + ..Default::default() + }; + + Box::pin(async move { + let response = utils::llm_api_request( + self.client + .post(MESSAGES_API_URL) + .header("anthropic-version", API_VERSION) + .header("content-type", "application/json") + .header("x-api-key", &self.api_key) + .json(&request), + "Anthropic", + ) + .await?; + + let stream = async_stream::stream! { + let mut sse_event_stream = utils::get_sse_events(response); + // let mut tool_calls = Vec::new(); + while let Some(event_result) = sse_event_stream.next().await { + match event_result { + Ok(event) => { + if let Some(chunk) = parse_anthropic_event(event) { + yield chunk; + } + }, + Err(e) => yield Err(e), + } + } + }; + + Ok(stream.boxed()) + }) + } +} diff --git a/server-new/src/llm/providers/anthropic/request.rs b/server-new/src/llm/providers/anthropic/request.rs new file mode 100644 index 0000000..cf95bee --- /dev/null +++ b/server-new/src/llm/providers/anthropic/request.rs @@ -0,0 +1,163 @@ +use serde::Serialize; + +use crate::llm::types::LlmMessage; + +pub fn build_anthropic_messages<'a>( + messages: &'a [LlmMessage], +) -> (Vec>, Option<&'a str>) { + let system_prompt = messages.iter().rev().find_map(|message| { + let LlmMessage::System(msg) = message else { + return None; + }; + Some(msg.as_str()) + }); + + let anthropic_messages: Vec = messages + .iter() + .filter_map(|message| { + let mut content_blocks = Vec::new(); + match message { + LlmMessage::User(user_message) => { + if !user_message.text.is_empty() { + content_blocks.push(AnthropicContentBlock::Text { + text: &user_message.text, + }); + } + // if let Some(ref files) = user_message.files { + // content_blocks.extend(files.iter().map(|file| match file.file_type { + // ChatRsFileType::Text => AnthropicContentBlock::Document { + // title: &file.name, + // source: AnthropicSource::Text { + // data: &file.content, + // media_type: "text/plain", + // }, + // }, + // ChatRsFileType::Image => AnthropicContentBlock::Image { + // source: AnthropicSource::Base64 { + // data: &file.content, + // media_type: &file.content_type, + // }, + // }, + // ChatRsFileType::Pdf => AnthropicContentBlock::Document { + // title: &file.name, + // source: AnthropicSource::Base64 { + // data: &file.content, + // media_type: "application/pdf", + // }, + // }, + // })); + // } + Some(AnthropicMessage { + role: "user", + content: content_blocks, + }) + } + LlmMessage::Assistant(assistant_message) => { + if !assistant_message.text.is_empty() { + content_blocks.push(AnthropicContentBlock::Text { + text: &assistant_message.text, + }); + } + // if let Some(ref tool_calls) = assistant_message.tool_calls { + // content_blocks.extend(tool_calls.iter().map(|tc| { + // AnthropicContentBlock::ToolUse { + // id: &tc.id, + // name: &tc.tool_name, + // input: &tc.parameters, + // } + // })); + // } + Some(AnthropicMessage { + role: "assistant", + content: content_blocks, + }) + } + // LlmMessage::Tool(result) => { + // content_blocks.push(AnthropicContentBlock::ToolResult { + // tool_use_id: &result.tool_call_id, + // content: &result.content, + // }); + // Some(AnthropicMessage { + // role: "user", + // content: content_blocks, + // }) + // } + _ => None, + } + }) + .collect(); + + (anthropic_messages, system_prompt) +} + +// pub fn build_anthropic_tools<'a>(tools: &'a [LlmTool]) -> Vec> { +// tools +// .iter() +// .map(|tool| AnthropicTool { +// name: &tool.name, +// description: &tool.description, +// input_schema: &tool.input_schema, +// }) +// .collect() +// } + +/// Anthropic API request message +#[derive(Debug, Serialize)] +pub struct AnthropicMessage<'a> { + pub role: &'a str, + pub content: Vec>, +} + +/// Anthropic API request body +#[derive(Debug, Default, Serialize)] +pub struct AnthropicRequest<'a> { + pub model: &'a str, + pub messages: Vec>, + pub max_tokens: u32, + #[serde(skip_serializing_if = "Option::is_none")] + pub temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub system: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + pub stream: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub tools: Option>>, +} + +/// Anthropic tool definition +#[derive(Debug, Serialize)] +pub struct AnthropicTool<'a> { + name: &'a str, + description: &'a str, + input_schema: &'a serde_json::Value, +} + +/// Anthropic content block for messages +#[derive(Debug, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum AnthropicContentBlock<'a> { + Text { text: &'a str }, + // Image { + // source: AnthropicSource<'a>, + // }, + // Document { + // title: &'a str, + // source: AnthropicSource<'a>, + // }, + // ToolUse { + // id: &'a str, + // name: &'a str, + // input: &'a HashMap, + // }, + // ToolResult { + // tool_use_id: &'a str, + // content: &'a str, + // }, +} + +// #[derive(Debug, Serialize)] +// #[serde(tag = "type", rename_all = "lowercase")] +// pub enum AnthropicSource<'a> { +// Base64 { data: &'a str, media_type: &'a str }, +// Text { data: &'a str, media_type: &'a str }, +// } diff --git a/server-new/src/llm/providers/anthropic/response.rs b/server-new/src/llm/providers/anthropic/response.rs new file mode 100644 index 0000000..b71568e --- /dev/null +++ b/server-new/src/llm/providers/anthropic/response.rs @@ -0,0 +1,205 @@ +use serde::Deserialize; + +use crate::llm::{ + error::LlmStreamChunkError, + interface::{LlmStreamChunk, LlmStreamChunkResult}, + types::LlmUsage, +}; + +/// Parse an Anthropic SSE event. +pub fn parse_anthropic_event( + event: AnthropicStreamEvent, + // tools: Option<&Vec>, + // tool_calls: &mut Vec, +) -> Option { + match event { + AnthropicStreamEvent::MessageStart { message } => { + if let Some(usage) = message.usage { + return Some(Ok(LlmStreamChunk::Usage(usage.into()))); + } + } + AnthropicStreamEvent::ContentBlockStart { content_block, .. } => match content_block { + AnthropicResponseContentBlock::Text { text } => { + return Some(Ok(LlmStreamChunk::Text(text))); + } + AnthropicResponseContentBlock::ToolUse { .. } => { + // tool_calls.push(AnthropicStreamToolCall { + // id, + // index, + // name, + // input: String::with_capacity(100), + // }); + } + }, + AnthropicStreamEvent::ContentBlockDelta { delta, .. } => match delta { + AnthropicDelta::TextDelta { text } => { + return Some(Ok(LlmStreamChunk::Text(text))); + } + AnthropicDelta::InputJsonDelta { .. } => { + // if let Some(tool_call) = tool_calls.iter_mut().find(|tc| tc.index == index) { + // tool_call.input.push_str(&partial_json); + // let chunk = LlmStreamChunk::PendingToolCall(LlmPendingToolCall { + // index, + // tool_name: tool_call.name.clone(), + // }); + // return Some(Ok(chunk)); + // } + } + }, + AnthropicStreamEvent::ContentBlockStop { .. } => { + // if let Some(llm_tools) = tools { + // if let Some(tc) = tool_calls + // .iter() + // .position(|tc| tc.index == index) + // .map(|i| tool_calls.swap_remove(i)) + // { + // if let Some(tool_call) = tc.convert(llm_tools) { + // let chunk = LlmStreamChunk::ToolCalls(vec![tool_call]); + // return Some(Ok(chunk)); + // } + // } + // } + } + AnthropicStreamEvent::MessageDelta { usage } => { + if let Some(usage) = usage { + return Some(Ok(LlmStreamChunk::Usage(usage.into()))); + } + } + AnthropicStreamEvent::Error { error } => { + let error_msg = format!("{}: {}", error.error_type, error.message); + return Some(Err(LlmStreamChunkError::Provider(error_msg))); + } + _ => {} // Ignore other events (ping, message_stop) + } + None +} + +/// Anthropic API response content block +#[derive(Debug, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum AnthropicResponseContentBlock { + Text { text: String }, + ToolUse { id: String, name: String }, +} + +/// Anthropic API response usage +#[derive(Debug, Deserialize)] +pub struct AnthropicUsage { + input_tokens: Option, + output_tokens: Option, +} + +impl From for LlmUsage { + fn from(usage: AnthropicUsage) -> Self { + LlmUsage { + input_tokens: usage.input_tokens, + output_tokens: usage.output_tokens, + cost: None, + } + } +} + +/// Anthropic API response +#[derive(Debug, Deserialize)] +pub struct AnthropicResponse { + pub content: Vec, + pub usage: Option, +} + +/// Anthropic stream response (message start) +#[derive(Debug, Deserialize)] +pub struct AnthropicStreamResponse { + // id: String, + // #[serde(rename = "type")] + // message_type: String, + // role: String, + // content: Vec, + // model: String, + // stop_reason: Option, + // stop_sequence: Option, + usage: Option, +} + +/// Anthropic streaming event types +#[derive(Debug, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum AnthropicStreamEvent { + MessageStart { + message: AnthropicStreamResponse, + }, + ContentBlockStart { + index: usize, + content_block: AnthropicResponseContentBlock, + }, + ContentBlockDelta { + index: usize, + delta: AnthropicDelta, + }, + ContentBlockStop { + index: usize, + }, + MessageDelta { + // delta: AnthropicMessageDelta, + usage: Option, + }, + MessageStop, + Ping, + Error { + error: AnthropicError, + }, +} + +#[derive(Debug, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum AnthropicDelta { + TextDelta { text: String }, + InputJsonDelta { partial_json: String }, +} + +// #[derive(Debug, Deserialize)] +// pub struct AnthropicMessageDelta { +// stop_reason: Option, +// stop_sequence: Option, +// } + +#[derive(Debug, Deserialize)] +pub struct AnthropicError { + #[serde(rename = "type")] + error_type: String, + message: String, +} + +// /// Helper struct for tracking streaming tool calls +// #[derive(Debug)] +// pub struct AnthropicStreamToolCall { +// /// Anthropic tool call ID +// id: String, +// /// Index of the tool call in the message +// index: usize, +// /// Name of the tool +// name: String, +// /// Partial input parameters (JSON stringified) +// input: String, +// } + +// impl AnthropicStreamToolCall { +// /// Convert Anthropic tool call format to ChatRsToolCall +// fn convert(self, llm_tools: &[LlmTool]) -> Option { +// let input = if self.input.trim().is_empty() { +// "{}" +// } else { +// &self.input +// }; +// let parameters = serde_json::from_str(input).ok()?; +// llm_tools +// .iter() +// .find(|tool| tool.name == self.name) +// .map(|tool| ChatRsToolCall { +// id: self.id, +// tool_id: tool.tool_id, +// tool_name: self.name, +// tool_type: tool.tool_type, +// parameters, +// }) +// } +// } diff --git a/server-new/src/llm/providers/mod.rs b/server-new/src/llm/providers/mod.rs index 06f37fa..10c2227 100644 --- a/server-new/src/llm/providers/mod.rs +++ b/server-new/src/llm/providers/mod.rs @@ -1,6 +1,10 @@ +mod anthropic; mod lorem; +mod ollama; mod openai; mod utils; +pub use anthropic::AnthropicProvider; pub use lorem::LoremProvider; +pub use ollama::OllamaProvider; pub use openai::{OpenAIProvider, OpenAIProviderConfig}; diff --git a/server-new/src/llm/providers/ollama/mod.rs b/server-new/src/llm/providers/ollama/mod.rs new file mode 100644 index 0000000..c69ae51 --- /dev/null +++ b/server-new/src/llm/providers/ollama/mod.rs @@ -0,0 +1,126 @@ +//! Ollama LLM provider + +use futures::StreamExt; + +use crate::llm::{ + error::LlmRequestError, + interface::{LlmPromptResponse, LlmProvider, LlmStreamingResponse}, + providers::utils, + types::{LlmChatRequest, LlmPrompt}, +}; + +mod request; +mod response; + +use {request::*, response::*}; + +const CHAT_API_URL: &str = "/api/chat"; +const COMPLETION_API_URL: &str = "/api/generate"; + +/// Ollama chat provider +#[derive(Debug, Clone)] +pub struct OllamaProvider { + client: reqwest::Client, + base_url: String, +} + +impl OllamaProvider { + pub fn new(http_client: &reqwest::Client, base_url: &str) -> Self { + Self { + client: http_client.clone(), + base_url: base_url.trim_end_matches('/').to_string(), + } + } +} + +impl LlmProvider for OllamaProvider { + fn prompt<'r>(&'r self, prompt: LlmPrompt<'r>) -> LlmPromptResponse<'r> { + let ollama_options = OllamaOptions { + temperature: prompt.options.temperature, + num_predict: prompt.options.max_tokens, + ..Default::default() + }; + let request = OllamaCompletionRequest { + model: &prompt.options.model, + prompt: prompt.text, + stream: Some(false), + options: Some(ollama_options), + }; + + Box::pin(async move { + let res: OllamaCompletionResponse = utils::llm_api_request( + self.client + .post(format!("{}{}", self.base_url, COMPLETION_API_URL)) + .header("content-type", "application/json") + .json(&request), + "Ollama", + ) + .await? + .json() + .await?; + + if let Some(usage) = res.usage() { + tracing::info!("Prompt usage: {:?}", usage); + } + if res.response.is_empty() { + return Err(LlmRequestError::NoContent); + } + + Ok(res.response) + }) + } + + fn stream_chat<'r>(&'r self, req: LlmChatRequest<'r>) -> LlmStreamingResponse<'r> { + let ollama_messages = build_ollama_messages(&req.messages); + // let ollama_tools = tools.as_ref().map(|t| build_ollama_tools(t)); + let ollama_options = OllamaOptions { + temperature: req.options.temperature, + num_predict: req.options.max_tokens, + ..Default::default() + }; + let request = OllamaChatRequest { + model: &req.options.model, + messages: ollama_messages, + // tools: ollama_tools, + stream: Some(true), + options: Some(ollama_options), + ..Default::default() + }; + + Box::pin(async move { + let response = utils::llm_api_request( + self.client + .post(format!("{}{}", self.base_url, CHAT_API_URL)) + .header("content-type", "application/json") + .json(&request), + "Ollama", + ) + .await?; + let stream = async_stream::stream! { + let mut json_stream = utils::get_json_events(response); + // let mut tool_calls: Vec = Vec::new(); + while let Some(event) = json_stream.next().await { + match event { + Ok(event) => { + for chunk in parse_ollama_event(event) { + yield Ok(chunk); + } + } + Err(e) => yield Err(e), + } + } + // if !tool_calls.is_empty() { + // if let Some(llm_tools) = tools { + // let converted = tool_calls + // .into_iter() + // .filter_map(|tc| tc.function.convert(&llm_tools)) + // .collect(); + // yield Ok(LlmStreamChunk::ToolCalls(converted)); + // } + // } + }; + + Ok(stream.boxed()) + }) + } +} diff --git a/server-new/src/llm/providers/ollama/request.rs b/server-new/src/llm/providers/ollama/request.rs new file mode 100644 index 0000000..5ccc349 --- /dev/null +++ b/server-new/src/llm/providers/ollama/request.rs @@ -0,0 +1,148 @@ +//! Ollama API request structures + +use serde::Serialize; +use serde_with::skip_serializing_none; + +use crate::llm::types::LlmMessage; + +/// Convert LlmMessages to Ollama messages +pub fn build_ollama_messages(messages: &[LlmMessage]) -> Vec> { + messages + .iter() + .map(|message| match message { + LlmMessage::User(user_message) => { + // let images = user_message.files.as_ref().map(|files| { + // files + // .iter() + // .filter_map(|file| match file.file_type { + // ChatRsFileType::Image => Some(file.content.as_str()), + // _ => None, + // }) + // .collect::>() + // }); + OllamaMessage { + role: "user", + content: &user_message.text, + // images, + ..Default::default() + } + } + LlmMessage::Assistant(assistant_message) => { + // let tool_calls = assistant_message.tool_calls.as_ref().map(|tool_calls| { + // tool_calls + // .iter() + // .map(|tc| OllamaToolCall { + // function: OllamaFunction { + // name: &tc.tool_name, + // arguments: &tc.parameters, + // }, + // }) + // .collect() + // }); + OllamaMessage { + role: "assistant", + content: &assistant_message.text, + // tool_calls, + ..Default::default() + } + } + LlmMessage::System(text) => OllamaMessage { + role: "system", + content: text, + ..Default::default() + }, + // LlmMessage::Tool(result) => OllamaMessage { + // role: "tool", + // content: &result.content, + // tool_name: Some(&result.tool_name), + // ..Default::default() + // }, + }) + .collect() +} + +// /// Convert LlmTools to Ollama tools +// pub fn build_ollama_tools(tools: &[LlmTool]) -> Vec> { +// tools +// .iter() +// .map(|tool| OllamaTool { +// r#type: "function", +// function: OllamaToolSpec { +// name: &tool.name, +// description: &tool.description, +// parameters: &tool.input_schema, +// }, +// }) +// .collect() +// } + +/// Ollama chat request structure +#[skip_serializing_none] +#[derive(Debug, Default, Serialize)] +pub struct OllamaChatRequest<'a> { + pub model: &'a str, + pub messages: Vec>, + // pub tools: Option>>, + pub stream: Option, + pub options: Option, +} + +/// Ollama completion request structure +#[skip_serializing_none] +#[derive(Debug, Serialize)] +pub struct OllamaCompletionRequest<'a> { + pub model: &'a str, + pub prompt: &'a str, + pub stream: Option, + pub options: Option, +} + +/// Ollama chat message +#[skip_serializing_none] +#[derive(Debug, Default, Serialize)] +pub struct OllamaMessage<'a> { + pub role: &'a str, + pub content: &'a str, + pub images: Option>, + // pub tool_calls: Option>>, + pub tool_name: Option<&'a str>, +} + +// /// Ollama tool call in a message +// #[derive(Debug, Serialize)] +// pub struct OllamaToolCall<'a> { +// pub function: OllamaFunction<'a>, +// } + +// /// Ollama tool function +// #[derive(Debug, Serialize)] +// pub struct OllamaFunction<'a> { +// pub name: &'a str, +// pub arguments: &'a ToolParameters, +// } + +// /// Ollama tool definition +// #[derive(Debug, Serialize)] +// pub struct OllamaTool<'a> { +// pub r#type: &'a str, +// pub function: OllamaToolSpec<'a>, +// } + +// /// Ollama tool specification +// #[derive(Debug, Serialize)] +// pub struct OllamaToolSpec<'a> { +// pub name: &'a str, +// pub description: &'a str, +// pub parameters: &'a serde_json::Value, +// } + +/// Ollama model options +#[skip_serializing_none] +#[derive(Debug, Default, Serialize)] +pub struct OllamaOptions { + pub temperature: Option, + pub num_predict: Option, // Ollama's equivalent to max_tokens + pub top_p: Option, + pub top_k: Option, + pub seed: Option, +} diff --git a/server-new/src/llm/providers/ollama/response.rs b/server-new/src/llm/providers/ollama/response.rs new file mode 100644 index 0000000..dd1562f --- /dev/null +++ b/server-new/src/llm/providers/ollama/response.rs @@ -0,0 +1,148 @@ +//! Ollama API response structures + +use serde::Deserialize; + +use crate::llm::{interface::LlmStreamChunk, types::LlmUsage}; + +/// Parse Ollama streaming event into LlmStreamChunks +pub fn parse_ollama_event( + event: OllamaStreamEvent, + // tool_calls: &mut Vec, +) -> impl Iterator { + // Handle usage stats + let usage = event.usage().map(LlmStreamChunk::Usage); + + // Handle text response + let text = (!event.message.content.is_empty()) + .then_some(event.message.content) + .map(LlmStreamChunk::Text); + + [text, usage].into_iter().flatten() + + // Handle tool calls in the message + // if !event.message.tool_calls.is_empty() { + // for (index, tc) in event.message.tool_calls.iter().enumerate() { + // let tool_call = LlmPendingToolCall { + // index, + // tool_name: tc.function.name.clone(), + // }; + // chunks.push(Ok(LlmStreamChunk::PendingToolCall(tool_call))); + // } + // tool_calls.extend(event.message.tool_calls); + // } +} + +/// Ollama chat response (streaming) +#[derive(Debug, Deserialize)] +pub struct OllamaStreamEvent { + // pub model: String, + // pub created_at: String, + pub message: OllamaMessageResponse, + pub done: bool, + // #[serde(default)] + // pub done_reason: Option, + // #[serde(default)] + // pub total_duration: Option, + // #[serde(default)] + // pub load_duration: Option, + #[serde(default)] + pub prompt_eval_count: Option, + // #[serde(default)] + // pub prompt_eval_duration: Option, + #[serde(default)] + pub eval_count: Option, + // #[serde(default)] + // pub eval_duration: Option, +} + +/// Ollama completion response (non-streaming) +#[derive(Debug, Deserialize)] +pub struct OllamaCompletionResponse { + pub response: String, + // pub model: String, + // pub created_at: String, + // pub done: bool, + // #[serde(default)] + // pub done_reason: Option, + // #[serde(default)] + // pub total_duration: Option, + // #[serde(default)] + // pub load_duration: Option, + // #[serde(default)] + // pub prompt_eval_duration: Option, + // #[serde(default)] + // pub eval_duration: Option, + #[serde(default)] + pub prompt_eval_count: Option, + #[serde(default)] + pub eval_count: Option, +} + +/// Ollama message in response +#[derive(Debug, Deserialize)] +pub struct OllamaMessageResponse { + #[serde(default)] + pub content: String, + // pub role: String, + // #[serde(default)] + // pub tool_calls: Vec, +} + +// /// Ollama tool call in response +// #[derive(Debug, Deserialize)] +// pub struct OllamaToolCallResponse { +// pub function: OllamaFunctionResponse, +// } + +// /// Ollama tool function in response +// #[derive(Debug, Deserialize)] +// pub struct OllamaFunctionResponse { +// pub name: String, +// pub arguments: serde_json::Value, +// } + +// impl OllamaFunctionResponse { +// /// Convert to ChatRsToolCall if the tool exists in the provided tools +// pub fn convert(self, tools: &[LlmTool]) -> Option { +// let tool = tools.iter().find(|t| t.name == self.name)?; +// let parameters = serde_json::from_value(self.arguments).ok()?; + +// Some(ChatRsToolCall { +// id: uuid::Uuid::new_v4().to_string(), +// parameters, +// tool_id: tool.tool_id, +// tool_name: self.name, +// tool_type: tool.tool_type, +// }) +// } +// } + +impl OllamaCompletionResponse { + /// Convert usage to LlmUsage + pub fn usage(&self) -> Option { + if self.prompt_eval_count.is_some() || self.eval_count.is_some() { + Some(LlmUsage { + input_tokens: self.prompt_eval_count, + output_tokens: self.eval_count, + ..Default::default() + }) + } else { + None + } + } +} + +impl OllamaStreamEvent { + /// If last event in stream, convert usage to LlmUsage + pub fn usage(&self) -> Option { + if self.done && (self.prompt_eval_count.is_some() || self.eval_count.is_some()) { + Some(LlmUsage { + input_tokens: self.prompt_eval_count, + output_tokens: self.eval_count, + ..Default::default() + }) + } else { + None + } + } +} diff --git a/server-new/src/llm/providers/openai/mod.rs b/server-new/src/llm/providers/openai/mod.rs index 1577ef7..298b88c 100644 --- a/server-new/src/llm/providers/openai/mod.rs +++ b/server-new/src/llm/providers/openai/mod.rs @@ -140,15 +140,16 @@ impl LlmProvider for OpenAIProvider { Box::pin(async move { let provider_name = self.config.subtype.name(); - let response = utils::llm_api_request( - &self.client, + let mut response: OpenAIResponse = utils::llm_api_request( + self.client + .post(&format!("{}/chat/completions", self.config.base_url)) + .bearer_auth(&self.config.api_key) + .json(&request), provider_name, - &format!("{}/chat/completions", self.config.base_url), - &self.config.api_key, - &request, ) + .await? + .json() .await?; - let mut response: OpenAIResponse = response.json().await?; let text = response .choices @@ -187,11 +188,11 @@ impl LlmProvider for OpenAIProvider { Box::pin(async move { let response = utils::llm_api_request( - &self.client, + self.client + .post(&format!("{}/chat/completions", self.config.base_url)) + .bearer_auth(&self.config.api_key) + .json(&request), provider_name, - &format!("{}/chat/completions", self.config.base_url), - &self.config.api_key, - &request, ) .await?; diff --git a/server-new/src/llm/providers/utils.rs b/server-new/src/llm/providers/utils.rs index 23b32a7..514dbc5 100644 --- a/server-new/src/llm/providers/utils.rs +++ b/server-new/src/llm/providers/utils.rs @@ -1,7 +1,7 @@ //! Utilities for working with LLM requests and responses use futures::TryStreamExt; -use serde::{Serialize, de::DeserializeOwned}; +use serde::de::DeserializeOwned; use tokio_stream::{Stream, StreamExt}; use tokio_util::{ codec::{FramedRead, LinesCodec}, @@ -44,7 +44,7 @@ pub fn get_sse_events( }) } -/// Get a stream of deserialized events from a provider JSON stream, not SSE (e.g. Ollama uses this format). +/// Get a stream of deserialized events from a provider JSON Lines stream (e.g. Ollama uses this format). pub fn get_json_events( response: reqwest::Response, ) -> impl Stream> { @@ -57,17 +57,11 @@ pub fn get_json_events( } /// Convenience function to make an API request to an LLM provider -pub async fn llm_api_request( - client: &reqwest::Client, +pub async fn llm_api_request( + request: reqwest::RequestBuilder, provider_name: &str, - url: &str, - token: &str, - request: &Req, ) -> Result { - let response = client - .post(url) - .bearer_auth(token) - .json(&request) + let response = request .send() .await .map_err(|e| LlmRequestError::Provider(format!("{provider_name} request failed: {e}")))?; diff --git a/server-new/src/services/provider/mod.rs b/server-new/src/services/provider/mod.rs index a82ef87..1082108 100644 --- a/server-new/src/services/provider/mod.rs +++ b/server-new/src/services/provider/mod.rs @@ -10,10 +10,7 @@ use crate::{ OpenAISubtype, UpdateChatRsProvider, UpdateChatRsSecret, }, }, - llm::{ - interface::LlmProvider, - providers::{LoremProvider, OpenAIProvider, OpenAIProviderConfig}, - }, + llm::{interface::LlmProvider, providers::*}, services::{ auth::encryption::Encryptor, provider::{ @@ -66,16 +63,17 @@ impl<'r> ProviderService<'r> { provider.base_url, ), )), - _ => todo!(), - // ChatRsProviderType::Anthropic => Box::new(AnthropicProvider::new( - // http_client, - // redis, - // api_key.ok_or(ProviderError::MissingApiKey)?, - // )), - // ChatRsProviderType::Ollama => Box::new(OllamaProvider::new( - // http_client, - // base_url.unwrap_or("http://localhost:11434"), - // )), + ChatRsProviderType::Anthropic => Arc::new(AnthropicProvider::new( + self.http_client, + api_key.ok_or(ProviderError::MissingApiKey)?, + )), + ChatRsProviderType::Ollama => Arc::new(OllamaProvider::new( + self.http_client, + provider + .base_url + .as_deref() + .unwrap_or("http://localhost:11434"), + )), }; Ok(llm_provider) From 8a9635aeb04bd046435bdb9cb3b3e97b898eef6f Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Tue, 14 Jul 2026 21:32:25 -0400 Subject: [PATCH 084/111] add website files --- server-new/Cargo.lock | 29 +++++++++++++++++++++++++ server-new/Cargo.toml | 2 +- server-new/Dockerfile | 3 ++- server-new/src/config.rs | 2 ++ server-new/src/lib.rs | 1 + server-new/src/plugins/mod.rs | 1 + server-new/src/plugins/web.rs | 41 +++++++++++++++++++++++++++++++++++ 7 files changed, 77 insertions(+), 2 deletions(-) create mode 100644 server-new/src/plugins/web.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index c5ed6a3..31d595a 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -1279,6 +1279,12 @@ dependencies = [ "pin-project-lite", ] +[[package]] +name = "http-range-header" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9171a2ea8a68358193d15dd5d70c1c10a2afc3e7e4c5bc92bc9f025cebd7359c" + [[package]] name = "httparse" version = "1.10.1" @@ -1716,6 +1722,16 @@ version = "0.3.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a" +[[package]] +name = "mime_guess" +version = "2.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f7c44f8e672c00fe5308fa235f821cb4198414e1c77935c1ab6948d3fd78550e" +dependencies = [ + "mime", + "unicase", +] + [[package]] name = "minimal-lexical" version = "0.2.1" @@ -3274,12 +3290,19 @@ checksum = "b11f75e912b0c2be01b63d8cf8057b8c3f97cf34abb3d431a3a4c8675498e233" dependencies = [ "bitflags", "bytes", + "futures-core", + "futures-util", "http", "http-body", "http-body-util", + "http-range-header", + "httpdate", + "mime", + "mime_guess", "percent-encoding", "pin-project-lite", "tokio", + "tokio-util", "tower-layer", "tower-service", "tracing", @@ -3497,6 +3520,12 @@ dependencies = [ "version_check", ] +[[package]] +name = "unicase" +version = "2.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dbc4bc3a9f746d862c45cb89d705aa10f187bb96c76001afab07a0d35ce60142" + [[package]] name = "unicode-bidi" version = "0.3.18" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 8224ef2..cebc02a 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -89,7 +89,7 @@ tokio-util = { version = "0.7.18", features = ["io"] } tower = { version = "0.5", default-features = false } tower-http = { version = "0.7.0", - features = ["limit", "request-id", "timeout", "trace"] + features = ["fs", "limit", "request-id", "timeout", "trace"] } tower-sessions = { version = "0.15.0", diff --git a/server-new/Dockerfile b/server-new/Dockerfile index fe58b19..7c6c8e5 100644 --- a/server-new/Dockerfile +++ b/server-new/Dockerfile @@ -1,4 +1,4 @@ -# Image versions (can be overridden by args when building) +# Image versions ARG RUST_VERSION=1.96 ARG DEBIAN_VERSION=bookworm @@ -35,4 +35,5 @@ COPY --from=build /app/run-server /usr/local/bin/ # Run server WORKDIR /app ENV RS_CHAT_SERVER__HOST=0.0.0.0 +ENV RS_CHAT_SERVER__STATIC_PATH=/var/www CMD ["run-server"] diff --git a/server-new/src/config.rs b/server-new/src/config.rs index 8a61818..af304a4 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -36,6 +36,7 @@ pub struct ServerConfig { pub port: u16, pub base_url: String, pub log_level: String, + pub web_root: String, pub request_id_header: String, pub ip_header: Option, } @@ -46,6 +47,7 @@ impl Default for ServerConfig { port: 8080, base_url: String::from("http://localhost:8080"), log_level: String::from("info"), + web_root: String::from("../web/dist"), request_id_header: String::from("x-request-id"), ip_header: None, } diff --git a/server-new/src/lib.rs b/server-new/src/lib.rs index 138718a..66f6e63 100644 --- a/server-new/src/lib.rs +++ b/server-new/src/lib.rs @@ -31,6 +31,7 @@ pub async fn create_app() -> anyhow::Result> .register(api::plugin()) // Add API routes .register(plugins::auth::plugin()) // Setup auth & sessions .register(plugins::logging::plugin()) // Request logging + .register(plugins::web::plugin()) // Web app .register(plugins::security::plugin()) // Body limit, security headers, etc. .init() .await?; diff --git a/server-new/src/plugins/mod.rs b/server-new/src/plugins/mod.rs index fe892b2..5c3e338 100644 --- a/server-new/src/plugins/mod.rs +++ b/server-new/src/plugins/mod.rs @@ -8,6 +8,7 @@ pub mod database; pub mod logging; pub mod redis; pub mod security; +pub mod web; /// Shared plugin type with correct state and config type parameters pub type AxumPlugin = AdHocPlugin; diff --git a/server-new/src/plugins/web.rs b/server-new/src/plugins/web.rs new file mode 100644 index 0000000..1e9b125 --- /dev/null +++ b/server-new/src/plugins/web.rs @@ -0,0 +1,41 @@ +use axum::{ + Router, + http::{HeaderValue, header}, + middleware, + response::Response, +}; +use tower_http::services::{ServeDir, ServeFile}; + +use crate::plugins::AxumPlugin; + +/// Adds the website / static files to the router +pub fn plugin() -> AxumPlugin { + AxumPlugin::named("Web").on_setup(|app, router| { + let web_files_root = &app.config().server.web_root; + let web_service = ServeDir::new(web_files_root) + .fallback(ServeFile::new(format!("{web_files_root}/index.html"))); + let web_router = Router::new() + .fallback_service(web_service) + .layer(middleware::from_fn(cache_immutable_assets)); + + Ok(router.merge(web_router)) + }) +} + +/// Set cache headers for the website's immutable assets at `/assets/*` +async fn cache_immutable_assets( + req: axum::extract::Request, + next: axum::middleware::Next, +) -> Response { + let is_immutable_asset = req.uri().path().starts_with("/assets/"); + + let mut response = next.run(req).await; + if response.status().is_success() && is_immutable_asset { + response.headers_mut().insert( + header::CACHE_CONTROL, + HeaderValue::from_static("public, max-age=31536000, immutable"), + ); + } + + response +} From 35dcc9200f08e58f8f1dc5a56f50f67c5aa89780 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 01:27:33 -0400 Subject: [PATCH 085/111] add more chat/stream api routes --- server-new/src/api/chat.rs | 79 +++++++++++++++--- server-new/src/db/repositories/chat.rs | 35 -------- server-new/src/services/chat/error.rs | 5 +- server-new/src/services/chat/mod.rs | 108 +++++++++++++++++++------ server-new/src/services/chat/titles.rs | 6 +- server-new/src/services/stream/mod.rs | 44 ++++++++-- web/package.json | 2 +- web/src/components/Sidebar.tsx | 2 +- web/src/lib/api/client.ts | 2 +- web/vite.config.ts | 4 +- 10 files changed, 204 insertions(+), 83 deletions(-) diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index 0284746..80783c2 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -21,14 +21,41 @@ api_routes! { POST "/prompt" => prompt, "Prompt", { description: "Send a single prompt to a provider and get the response" }; - POST "/session/{session_id}" => chat_stream, "Chat", { + GET "/sessions" => get_active_streams, "Get active chat streams", { + description: "Get the session IDs that have ongoing response streams" + }; + GET "/sessions/{session_id}" => connect_chat_stream, "Access active chat stream", { + description: "Get a URL and token to access the response stream for this session" + }; + POST "/sessions/{session_id}" => chat_stream, "Stream chat", { description: "Send a message in a chat session and stream the response" }; - POST "/session/{session_id}/regenerate" => regenerate_response, "Regenerate response", { + POST "/sessions/{session_id}/cancel" => cancel_chat_stream, "Cancel active stream", { + description: "Cancel an ongoing chat stream" + }; + POST "/session/{session_id}/regenerate" => regenerate_response, "Regenerate chat response", { description: "Regenerate the latest assistant response in a chat session" }; } +async fn get_active_streams( + CurrentUser { user_id }: CurrentUser, + State(state): State, +) -> AppResult> { + let sessions = state + .chat_service() + .active_stream_sessions(&user_id) + .await?; + + Ok(Json(ActiveStreamsResponse { sessions })) +} + +#[derive(Debug, JsonSchema, serde::Serialize)] +struct ActiveStreamsResponse { + /// The chat session IDs that have ongoing response streams + sessions: Vec, +} + #[derive(Debug, Deserialize, JsonSchema)] struct PromptInput { /// The prompt to send to the LLM provider @@ -44,14 +71,15 @@ async fn prompt( Database(mut db): Database, State(state): State, Json(input): Json, -) -> AppResult> { +) -> AppResult> { let llm_provider = state .provider_service() .build_llm_provider(&mut db, &user_id, input.provider_id) .await?; - let text = state + let stream_access = state .chat_service() .prompt( + user_id, llm_provider, LlmUserMessage { text: input.message, @@ -61,12 +89,10 @@ async fn prompt( ) .await?; - Ok(Json(PromptResponse { text })) -} - -#[derive(Serialize, JsonSchema)] -struct PromptResponse { - text: String, + Ok(Json(StreamAccess { + url: stream_access.sse_url, + token: stream_access.token, + })) } #[derive(Debug, Deserialize, JsonSchema)] @@ -148,9 +174,40 @@ async fn regenerate_response( })) } -/// The URL and Bearer token to access the SSE stream +async fn connect_chat_stream( + CurrentUser { user_id }: CurrentUser, + Path(session_id): Path, + State(state): State, +) -> AppResult> { + let stream_access = state + .chat_service() + .connect_stream(&user_id, &session_id) + .await?; + + Ok(Json(StreamAccess { + url: stream_access.sse_url, + token: stream_access.token, + })) +} + +pub async fn cancel_chat_stream( + CurrentUser { user_id }: CurrentUser, + Path(session_id): Path, + State(state): State, +) -> AppResult<()> { + state + .chat_service() + .cancel_stream(&user_id, &session_id) + .await?; + + Ok(()) +} + +/// Access to an active streaming response #[derive(Serialize, JsonSchema)] struct StreamAccess { + /// URL to access the SSE stream url: String, + /// Bearer token to access the SSE stream token: String, } diff --git a/server-new/src/db/repositories/chat.rs b/server-new/src/db/repositories/chat.rs index e2f6300..f61d73f 100644 --- a/server-new/src/db/repositories/chat.rs +++ b/server-new/src/db/repositories/chat.rs @@ -142,41 +142,6 @@ impl<'a> ChatRepository<'a> { Ok((session.optional()?, messages?)) } - pub async fn find_messages_before( - &mut self, - user_id: &Uuid, - session_id: &Uuid, - message: &ChatRsMessage, - ) -> Result, diesel::result::Error> { - chat_messages::table - .inner_join(chat_sessions::table.on(chat_sessions::id.eq(chat_messages::session_id))) - .select(ChatRsMessage::as_select()) - .filter(chat_sessions::user_id.eq(user_id)) - .filter(chat_messages::session_id.eq(session_id)) - .filter(chat_messages::created_at.lt(message.created_at)) - .order_by(chat_messages::created_at.asc()) - .load(self.db) - .await - } - - pub async fn has_messages_after( - &mut self, - user_id: &Uuid, - session_id: &Uuid, - message: &ChatRsMessage, - ) -> Result { - let count = chat_messages::table - .inner_join(chat_sessions::table.on(chat_sessions::id.eq(chat_messages::session_id))) - .filter(chat_sessions::user_id.eq(user_id)) - .filter(chat_messages::session_id.eq(session_id)) - .filter(chat_messages::created_at.gt(message.created_at)) - .count() - .get_result::(self.db) - .await?; - - Ok(count > 0) - } - pub async fn search_sessions( &mut self, user_id: &Uuid, diff --git a/server-new/src/services/chat/error.rs b/server-new/src/services/chat/error.rs index ad2510c..b4367e8 100644 --- a/server-new/src/services/chat/error.rs +++ b/server-new/src/services/chat/error.rs @@ -7,6 +7,8 @@ pub enum ChatError { Messages, #[error("session not found")] SessionNotFound, + #[error("stream not found")] + StreamNotFound, #[error("no assistant response")] NoAssistantResponse, #[error("already streaming a response")] @@ -24,8 +26,9 @@ pub enum ChatError { impl From for AppError { fn from(value: ChatError) -> Self { match value { - ChatError::Messages => Self::bad_request("invalid messages"), ChatError::SessionNotFound => Self::not_found("chat session not found"), + ChatError::StreamNotFound => Self::not_found("stream not found for this session"), + ChatError::Messages => Self::bad_request("invalid messages"), ChatError::NoAssistantResponse => Self::bad_request("no assistant response found"), ChatError::AlreadyStreaming => Self::bad_request("already streaming this chat session"), ChatError::Request(err) => Self::bad_request(err.to_string()), diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index 82ea3a6..4309765 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -1,6 +1,6 @@ use std::sync::Arc; -use tinistream_client::types::StreamAccessResponse; +use tinistream_client::types::{StreamAccessResponse, StreamStatus}; use uuid::Uuid; use crate::{ @@ -13,7 +13,7 @@ use crate::{ }, llm::{ interface::LlmProvider, - types::{LlmChatOptions, LlmChatRequest, LlmPrompt, LlmUserMessage}, + types::{LlmChatOptions, LlmChatRequest, LlmMessage, LlmUserMessage}, }, services::{ chat::error::ChatError, @@ -32,7 +32,7 @@ pub struct ChatService<'r> { tinistream: &'r TinistreamClient, } -/// Chat stream parameters for response generation +#[derive(Debug)] struct ChatStreamParams { user_id: Uuid, session_id: Uuid, @@ -49,20 +49,82 @@ impl<'r> ChatService<'r> { } } - /// Send a simple prompt to the LLM provider + /// Connect to an ongoing stream + pub async fn connect_stream( + &self, + user_id: &Uuid, + session_id: &Uuid, + ) -> Result { + let stream_key = StreamingService::chat_stream_key(user_id, session_id); + let streams = StreamingService::new(self.tinistream); + if !streams.exists_stream(&stream_key).await? { + return Err(ChatError::StreamNotFound); + } + + Ok(streams.access_stream(&stream_key).await?) + } + + /// Cancel an ongoing stream + pub async fn cancel_stream( + &self, + user_id: &Uuid, + session_id: &Uuid, + ) -> Result { + let stream_key = StreamingService::chat_stream_key(user_id, session_id); + let streams = StreamingService::new(self.tinistream); + if !streams.exists_stream(&stream_key).await? { + return Err(ChatError::StreamNotFound); + } + + Ok(streams.cancel_stream(&stream_key).await?) + } + + /// Get the currently streaming session IDs for the given user + pub async fn active_stream_sessions(&self, user_id: &Uuid) -> Result, ChatError> { + let prefix = StreamingService::chat_stream_prefix(user_id); + let session_ids = StreamingService::new(self.tinistream) + .active_streams(&prefix) + .await? + .iter() + .filter_map(|stream| StreamingService::session_id_from_stream_key(&stream.key, &prefix)) + .collect(); + + Ok(session_ids) + } + + /// Send a single prompt to the LLM provider and stream the response pub async fn prompt( &self, + user_id: Uuid, provider: Arc, prompt: LlmUserMessage, options: LlmChatOptions, - ) -> Result { - let llm_prompt = LlmPrompt { - text: &prompt.text, + ) -> Result { + let request = LlmChatRequest { + messages: &[LlmMessage::User(prompt)], options: &options, }; - Ok(provider.prompt(llm_prompt).await?) + let response_stream = provider.stream_chat(request).await?; + + let stream_key = StreamingService::prompt_key(&user_id); + let (stream_access, ws_writer, ws_reader) = StreamingService::new(self.tinistream) + .create_stream(&stream_key) + .await?; + + // Spawn thread to process LLM streaming response + let tinistream_client = self.tinistream.to_owned(); + tokio::spawn(async move { + let _ = StreamingService::process_stream(response_stream, ws_writer, ws_reader).await; + let _ = StreamingService::new(&tinistream_client) + .end_stream(&stream_key) + .await; + }); + + // Return the URL and token for the user to access the client stream + Ok(stream_access) } + /// Stream response to a user message in a session pub async fn stream_user_chat( &self, db: &mut DbService, @@ -116,6 +178,7 @@ impl<'r> ChatService<'r> { .await } + /// Regenerate the last assistant response in a chat session pub async fn regenerate_response( &self, db: &mut DbService, @@ -151,36 +214,35 @@ impl<'r> ChatService<'r> { .await } + /// Start the LLM response stream and return the access URL & token async fn start_assistant_stream( &self, provider: Arc, messages: Vec, params: ChatStreamParams, ) -> Result { + let streams = StreamingService::new(self.tinistream); let stream_key = StreamingService::chat_stream_key(¶ms.user_id, ¶ms.session_id); - let stream_service = StreamingService::new(self.tinistream); - if stream_service.exists_stream(&stream_key).await? { + if streams.exists_stream(&stream_key).await? { return Err(ChatError::AlreadyStreaming); } - let llm_messages = messages::build_llm_messages(messages)?; let response_stream = provider .stream_chat(LlmChatRequest { - messages: &llm_messages, + messages: &messages::build_llm_messages(messages)?, options: ¶ms.chat_options, }) .await?; + let (stream_access, ws_writer, ws_reader) = streams.create_stream(&stream_key).await?; - let (stream_access, ws_writer, ws_reader) = - stream_service.create_stream(&stream_key).await?; - + // Spawn thread to process and save LLM streaming response let db_pool = self.db_pool.to_owned(); let tinistream_client = self.tinistream.to_owned(); tokio::spawn(async move { - let response = + let output = StreamingService::process_stream(response_stream, ws_writer, ws_reader).await; - let stream_cancelled = response.cancelled; - if let Err(err) = Self::persist_response(response, params, db_pool).await { + let stream_cancelled = output.cancelled; + if let Err(err) = Self::persist_response(output, params, db_pool).await { tracing::error!("Failed to save assistant response: {err}"); } @@ -197,7 +259,7 @@ impl<'r> ChatService<'r> { /// Save response message and metadata to database async fn persist_response( - response: LlmStreamOutput, + output: LlmStreamOutput, params: ChatStreamParams, db_pool: DbPool, ) -> Result { @@ -207,15 +269,15 @@ impl<'r> ChatService<'r> { provider_options: Some(params.chat_options), // tool_calls: response.tool_calls, // files: image_ids, - usage: response.usage, - errors: response.errors, - partial: response.cancelled.then_some(true), + usage: output.usage, + errors: output.errors, + partial: output.cancelled.then_some(true), ..Default::default() }; let new_message = db .chats() .save_message(NewChatRsMessage { - content: &response.text.unwrap_or_default(), + content: &output.text.unwrap_or_default(), meta: ChatRsMessageMeta::new_assistant(assistant_meta), role: ChatRsMessageRole::Assistant, session_id: ¶ms.session_id, diff --git a/server-new/src/services/chat/titles.rs b/server-new/src/services/chat/titles.rs index 63f4c3b..5b8808f 100644 --- a/server-new/src/services/chat/titles.rs +++ b/server-new/src/services/chat/titles.rs @@ -33,7 +33,7 @@ pub fn generate_title( } const TITLE_PROMPT: &str = "This is the first message sent by a human in a chat session with an AI chatbot. \ - Please generate a short title for the chat session (3-7 words) in plain text, with no quotes or prefixes"; + Please generate a short title for the chat session (3-7 words) in plain text, with no quotes or prefixes."; const TITLE_PROMPT_TEMPERATURE: f32 = 0.7; const TITLE_PROMPT_MAX_TOKENS: u32 = 20; @@ -45,10 +45,10 @@ async fn generate( model: String, db_pool: DbPool, ) -> Result<(), ChatError> { - let message = format!("{TITLE_PROMPT}: \"{user_message}\""); + let prompt = format!("{TITLE_PROMPT}\n\n\"{user_message}\""); let title = provider .prompt(LlmPrompt { - text: &message, + text: &prompt, options: &LlmChatOptions { model, temperature: Some(TITLE_PROMPT_TEMPERATURE), diff --git a/server-new/src/services/stream/mod.rs b/server-new/src/services/stream/mod.rs index eb8cf06..71b5d7f 100644 --- a/server-new/src/services/stream/mod.rs +++ b/server-new/src/services/stream/mod.rs @@ -4,7 +4,7 @@ use futures::{ }; use reqwest_websocket::WebSocket; use tinistream::TinistreamClient; -use tinistream_client::types::{StreamAccessResponse, StreamStatus}; +use tinistream_client::types::{StreamAccessResponse, StreamInfo, StreamStatus}; use uuid::Uuid; pub mod error; @@ -42,21 +42,42 @@ impl<'r> StreamingService<'r> { Self { tinistream } } - /// Get the key of the chat stream in Redis for the given user and session ID + /// Get the Redis key of the chat stream for the given user and session ID pub fn chat_stream_key(user_id: &Uuid, session_id: &Uuid) -> String { format!("{}{}", Self::chat_stream_prefix(user_id), session_id) } - /// Get the key prefix for the user's chat streams in Redis + /// Get the Redis key prefix for the user's chat streams pub fn chat_stream_prefix(user_id: &Uuid) -> String { - format!("user:{}:chat:", user_id) + format!("user:{user_id}:chat:") } - /// Check for existing client stream + /// Generate a Redis key for a user's prompt + pub fn prompt_key(user_id: &Uuid) -> String { + format!("user:{user_id}:prompt:{}", Uuid::new_v4()) + } + + /// Extract the session ID from the user's stream key + pub fn session_id_from_stream_key(key: &str, key_prefix: &str) -> Option { + key.strip_prefix(&key_prefix) + .and_then(|session_id| Uuid::try_parse(session_id).ok()) + } + + /// Check for existing active client stream pub async fn exists_stream(&self, stream_key: &str) -> Result { Ok(self.tinistream.stream_exists(&stream_key).await?) } + /// Currently active streams with the given prefix + pub async fn active_streams(&self, prefix: &str) -> Result, StreamingError> { + let streams = self + .tinistream + .active_streams(&format!("{prefix}*",)) + .await?; + + Ok(streams) + } + /// Start the client stream, and return a WebSocket writer and reader for it pub async fn create_stream( &self, @@ -68,6 +89,14 @@ impl<'r> StreamingService<'r> { Ok((stream_access, writer, reader)) } + /// Get access to an ongoing client stream + pub async fn access_stream( + &self, + stream_key: &str, + ) -> Result { + Ok(self.tinistream.stream_connect(stream_key).await?) + } + /// Process and write the LLM response stream via the WebSocket connection, /// and return the accumulated response. pub async fn process_stream( @@ -84,4 +113,9 @@ impl<'r> StreamingService<'r> { pub async fn end_stream(&self, stream_key: &str) -> Result { Ok(self.tinistream.stream_end(stream_key).await?) } + + /// Signal stream cancellation + pub async fn cancel_stream(&self, stream_key: &str) -> Result { + Ok(self.tinistream.stream_cancel(stream_key).await?) + } } diff --git a/web/package.json b/web/package.json index f7b36fa..d8f2abe 100644 --- a/web/package.json +++ b/web/package.json @@ -13,7 +13,7 @@ "lint": "biome lint", "lint:ci": "biome ci", "format": "biome check --linter-enabled=false", - "gen-api": "pnpm dlx openapi-typescript http://localhost:8000/api/openapi.json -o src/lib/api/types.d.ts" + "gen-api": "pnpm dlx openapi-typescript http://localhost:8080/api/v1/docs/openapi.json -o src/lib/api/types.d.ts" }, "dependencies": { "@radix-ui/react-alert-dialog": "^1.1.15", diff --git a/web/src/components/Sidebar.tsx b/web/src/components/Sidebar.tsx index ed606ab..9d5828e 100644 --- a/web/src/components/Sidebar.tsx +++ b/web/src/components/Sidebar.tsx @@ -83,7 +83,7 @@ export function AppSidebar({ ); const onLogout = React.useCallback(async () => { - await fetch(`${import.meta.env.VITE_API_URL || ""}/api/auth/logout`, { + await fetch(`${import.meta.env.VITE_API_URL || "/api/v1"}/auth/logout`, { method: "POST", }); queryClient.invalidateQueries({ queryKey: ["user"] }); diff --git a/web/src/lib/api/client.ts b/web/src/lib/api/client.ts index 7efc425..4ff822c 100644 --- a/web/src/lib/api/client.ts +++ b/web/src/lib/api/client.ts @@ -3,7 +3,7 @@ import createClient from "openapi-fetch"; import type { paths } from "./types"; -export const API_URL: string = import.meta.env.VITE_API_URL || "/api"; +export const API_URL: string = import.meta.env.VITE_API_URL || "/api/v1"; export const client = createClient({ baseUrl: API_URL, diff --git a/web/vite.config.ts b/web/vite.config.ts index a18f384..63b1a4c 100644 --- a/web/vite.config.ts +++ b/web/vite.config.ts @@ -16,8 +16,8 @@ export default defineConfig({ ], server: { proxy: { - "/api": { - target: "http://localhost:8000", + "/api/v1": { + target: "http://localhost:8080", changeOrigin: true, secure: false, }, From cf297afc0b83d98928079507d3f1858f1405dfe9 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 05:01:32 -0400 Subject: [PATCH 086/111] add request logs --- server-new/Cargo.lock | 47 ++++++++ server-new/Cargo.toml | 6 +- .../down.sql | 1 + .../up.sql | 25 ++++ server-new/src/api/chat.rs | 2 + server-new/src/db/mod.rs | 3 + server-new/src/db/models.rs | 4 +- server-new/src/db/models/log.rs | 71 +++++++++++ server-new/src/db/repositories.rs | 2 + server-new/src/db/repositories/log.rs | 86 +++++++++++++ server-new/src/db/schema.rs | 25 ++++ server-new/src/llm/error.rs | 6 +- server-new/src/llm/interface.rs | 23 +++- server-new/src/llm/providers/anthropic/mod.rs | 26 ++-- server-new/src/llm/providers/lorem.rs | 23 ++-- server-new/src/llm/providers/ollama/mod.rs | 15 ++- server-new/src/llm/providers/openai/mod.rs | 33 +++-- server-new/src/llm/providers/utils.rs | 9 ++ server-new/src/llm/types.rs | 2 +- server-new/src/services/chat/mod.rs | 114 +++++++++++++++--- server-new/src/services/chat/titles.rs | 94 ++++++++++----- server-new/src/services/stream/mod.rs | 13 ++ 22 files changed, 532 insertions(+), 98 deletions(-) create mode 100644 server-new/migrations/2026-07-15-053539-0000_add_request_logs/down.sql create mode 100644 server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql create mode 100644 server-new/src/db/models/log.rs create mode 100644 server-new/src/db/repositories/log.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index 31d595a..ec59900 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -335,6 +335,21 @@ version = "0.22.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" +[[package]] +name = "bigdecimal" +version = "0.4.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4d6867f1565b3aad85681f1015055b087fcfd840d6aeee6eee7f2da317603695" +dependencies = [ + "autocfg", + "libm", + "num-bigint", + "num-integer", + "num-traits", + "serde", + "serde_json", +] + [[package]] name = "bitflags" version = "2.13.0" @@ -760,12 +775,16 @@ version = "2.3.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e54d1f576cd3a3460f212a4615fd12ce1b6303c095b79a44449ffbe627753dc1" dependencies = [ + "bigdecimal", "bitflags", "byteorder", "chrono", "diesel_derives", "downcast-rs", "itoa", + "num-bigint", + "num-integer", + "num-traits", "serde_json", "uuid", ] @@ -1628,6 +1647,12 @@ version = "0.2.186" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" +[[package]] +name = "libm" +version = "0.2.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981" + [[package]] name = "libredox" version = "0.1.18" @@ -1768,12 +1793,31 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "num-bigint" +version = "0.4.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c89e69e7e0f03bea5ef08013795c25018e101932225a656383bd384495ecc367" +dependencies = [ + "num-integer", + "num-traits", +] + [[package]] name = "num-conv" version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "521739c6d2bac4aa25192232afe6841231376b2b26d4d9fae5ecf8ca5772e441" +[[package]] +name = "num-integer" +version = "0.1.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7969661fd2958a5cb096e56c8e1ad0444ac2bbcd0061bd28660485a44879858f" +dependencies = [ + "num-traits", +] + [[package]] name = "num-traits" version = "0.2.19" @@ -2394,6 +2438,8 @@ dependencies = [ "axum-aide-macros", "axum-helmet", "axum-plugin", + "bigdecimal", + "bon", "chrono", "diesel", "diesel-async", @@ -2553,6 +2599,7 @@ version = "1.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a2b42f36aa1cd011945615b92222f6bf73c599a102a300334cd7f8dbeec726cc" dependencies = [ + "bigdecimal", "chrono", "dyn-clone", "indexmap", diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index cebc02a..b0e46c4 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -26,6 +26,8 @@ axum-plugin = { rev = "be17dc9aec", features = ["figment"] } +bigdecimal = { version = "0.4.10", features = ["serde-json"] } +bon = "3.9.3" chrono = { version = "0.4.45", default-features = false, @@ -34,7 +36,7 @@ chrono = { diesel = { version = "2.3.11", default-features = false, - features = ["chrono", "serde_json", "uuid"] + features = ["chrono", "numeric", "serde_json", "uuid"] } diesel-async = { version = "0.9.2", @@ -59,7 +61,7 @@ reqwest = { reqwest-websocket = { version = "0.6.0", features = ["json"] } schemars = { version = "1.2.1", - features = ["chrono04", "preserve_order", "uuid1"] + features = ["bigdecimal04", "chrono04", "preserve_order", "uuid1"] } serde = { version = "1.0.228", features = ["derive"] } serde_json = "1.0.150" diff --git a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/down.sql b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/down.sql new file mode 100644 index 0000000..f2aff2c --- /dev/null +++ b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/down.sql @@ -0,0 +1 @@ +DROP TABLE logs; diff --git a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql new file mode 100644 index 0000000..c5cf98e --- /dev/null +++ b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql @@ -0,0 +1,25 @@ +CREATE TABLE logs ( + id SERIAL PRIMARY KEY, + kind text NOT NULL, -- chat, title, prompt, image, audio, etc. + user_id uuid NOT NULL REFERENCES users (id) ON DELETE CASCADE, + provider_id integer REFERENCES providers (id) ON DELETE SET NULL, + session_id uuid REFERENCES chat_sessions (id) ON DELETE SET NULL, + message_id uuid REFERENCES chat_messages (id) ON DELETE SET NULL, + model text NOT NULL, + request_id text, -- provider request ID + input_tokens integer, + output_tokens integer, + cost numeric(12, 6), + status text NOT NULL, -- completed, failed, cancelled + error text, + started_at timestamptz NOT NULL DEFAULT now(), + completed_at timestamptz +); + +CREATE INDEX logs_user_id_started_at_idx ON logs (user_id, started_at DESC); + +CREATE INDEX logs_session_id_idx ON logs (session_id); + +CREATE INDEX logs_message_id_idx ON logs (message_id); + +CREATE INDEX logs_provider_id_started_at_idx ON logs (provider_id, started_at DESC); diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index 80783c2..600939a 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -79,7 +79,9 @@ async fn prompt( let stream_access = state .chat_service() .prompt( + &mut db, user_id, + input.provider_id, llm_provider, LlmUserMessage { text: input.message, diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs index d78ec01..a6def9b 100644 --- a/server-new/src/db/mod.rs +++ b/server-new/src/db/mod.rs @@ -59,6 +59,9 @@ impl DbService { pub fn chats(&mut self) -> repositories::ChatRepository<'_> { repositories::ChatRepository::new(&mut self.cxn) } + pub fn logs(&mut self) -> repositories::LogRepository<'_> { + repositories::LogRepository::new(&mut self.cxn) + } pub fn providers(&mut self) -> repositories::ProviderRepository<'_> { repositories::ProviderRepository::new(&mut self.cxn) } diff --git a/server-new/src/db/models.rs b/server-new/src/db/models.rs index 4e69f44..38ce052 100644 --- a/server-new/src/db/models.rs +++ b/server-new/src/db/models.rs @@ -5,15 +5,17 @@ use crate::db::schema; mod api_key; mod chat; // mod file; +mod log; mod provider; mod secret; -// mod tool; mod session; +// mod tool; mod user; pub use api_key::*; pub use chat::*; // pub use file::*; +pub use log::*; pub use provider::*; pub use secret::*; // pub use tool::*; diff --git a/server-new/src/db/models/log.rs b/server-new/src/db/models/log.rs new file mode 100644 index 0000000..c2635fa --- /dev/null +++ b/server-new/src/db/models/log.rs @@ -0,0 +1,71 @@ +use bigdecimal::BigDecimal; +use diesel::prelude::*; +use schemars::JsonSchema; +use serde::Serialize; +use strum::{AsRefStr, EnumString}; +use uuid::Uuid; + +use crate::db::{UtcDateTime, models::ChatRsUser}; + +#[derive(Identifiable, Associations, Queryable, Selectable, Serialize, JsonSchema)] +#[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] +#[diesel(table_name = super::schema::logs)] +pub struct ChatRsLog { + pub id: i32, + pub kind: String, + pub user_id: Uuid, + pub provider_id: Option, + pub session_id: Option, + pub message_id: Option, + pub model: String, + pub request_id: Option, + pub input_tokens: Option, + pub output_tokens: Option, + pub cost: Option, + pub status: String, + pub error: Option, + pub started_at: UtcDateTime, + pub completed_at: Option, +} + +#[derive(Debug, Clone, Copy, EnumString, AsRefStr)] +#[strum(serialize_all = "lowercase")] +pub enum ChatRsLogKind { + Chat, + Title, + Prompt, + Image, +} + +#[derive(Debug, Clone, Copy, EnumString, AsRefStr)] +#[strum(serialize_all = "lowercase")] +pub enum ChatRsLogStatus { + Started, + Completed, + Failed, + Cancelled, +} + +#[derive(Insertable)] +#[diesel(table_name = super::schema::logs)] +pub struct NewChatRsLog<'a> { + pub kind: &'a str, + pub user_id: &'a Uuid, + pub provider_id: i32, + pub session_id: Option<&'a Uuid>, + pub model: &'a str, + pub status: &'a str, +} + +#[derive(Default, AsChangeset)] +#[diesel(table_name = super::schema::logs)] +pub struct UpdateChatRsLog<'a> { + pub message_id: Option<&'a Uuid>, + pub request_id: Option<&'a str>, + pub input_tokens: Option, + pub output_tokens: Option, + pub cost: Option, + pub status: &'a str, + pub error: Option<&'a str>, + pub completed_at: UtcDateTime, +} diff --git a/server-new/src/db/repositories.rs b/server-new/src/db/repositories.rs index 8db7026..f592d53 100644 --- a/server-new/src/db/repositories.rs +++ b/server-new/src/db/repositories.rs @@ -2,6 +2,7 @@ mod api_key; mod chat; +mod log; mod provider; mod secret; mod session; @@ -9,6 +10,7 @@ mod user; pub use api_key::ApiKeyRepository; pub use chat::ChatRepository; +pub use log::LogRepository; pub use provider::ProviderRepository; pub use secret::SecretRepository; pub use session::SessionRepository; diff --git a/server-new/src/db/repositories/log.rs b/server-new/src/db/repositories/log.rs new file mode 100644 index 0000000..d53d869 --- /dev/null +++ b/server-new/src/db/repositories/log.rs @@ -0,0 +1,86 @@ +use bigdecimal::{BigDecimal, FromPrimitive}; +use bon::bon; +use chrono::Utc; +use diesel::prelude::*; +use diesel_async::RunQueryDsl; +use uuid::Uuid; + +use crate::{ + db::{ + DbConnection, + models::{ChatRsLogKind, ChatRsLogStatus, NewChatRsLog, UpdateChatRsLog}, + schema::logs, + }, + llm::types::LlmUsage, +}; + +pub struct LogRepository<'a> { + db: &'a mut DbConnection, +} + +#[bon] +impl<'a> LogRepository<'a> { + pub fn new(db: &'a mut DbConnection) -> Self { + LogRepository { db } + } + + /// Create a new LLM request log entry + #[builder(finish_fn = "build")] + pub async fn create( + &mut self, + user_id: &Uuid, + provider_id: i32, + model: &str, + kind: ChatRsLogKind, + session_id: Option<&Uuid>, + ) -> QueryResult { + let new_log = NewChatRsLog { + kind: kind.as_ref(), + user_id, + provider_id, + session_id, + model, + status: ChatRsLogStatus::Started.as_ref(), + }; + + diesel::insert_into(logs::table) + .values(new_log) + .returning(logs::id) + .get_result(self.db) + .await + } + + /// Complete a LLM request log entry + #[builder(finish_fn = "build")] + pub async fn complete( + &mut self, + id: i32, + message_id: Option<&Uuid>, + request_id: Option<&str>, + usage: Option<&LlmUsage>, + error: Option<&str>, + status: ChatRsLogStatus, + ) -> QueryResult { + let update_log = UpdateChatRsLog { + message_id, + request_id, + input_tokens: usage + .and_then(|u| u.input_tokens) + .and_then(|t| t.try_into().ok()), + output_tokens: usage + .and_then(|u| u.output_tokens) + .and_then(|t| t.try_into().ok()), + cost: usage.and_then(|u| u.cost.and_then(BigDecimal::from_f32)), + error, + status: status.as_ref(), + completed_at: Utc::now(), + }; + + diesel::update(logs::table) + .filter(logs::id.eq(id)) + .set(update_log) + .returning(logs::id) + .get_result(self.db) + .await + } +} diff --git a/server-new/src/db/schema.rs b/server-new/src/db/schema.rs index fbae7bc..9e6218b 100644 --- a/server-new/src/db/schema.rs +++ b/server-new/src/db/schema.rs @@ -84,6 +84,26 @@ diesel::table! { } } +diesel::table! { + logs (id) { + id -> Int4, + kind -> Text, + user_id -> Uuid, + provider_id -> Nullable, + session_id -> Nullable, + message_id -> Nullable, + model -> Text, + request_id -> Nullable, + input_tokens -> Nullable, + output_tokens -> Nullable, + cost -> Nullable, + status -> Text, + error -> Nullable, + started_at -> Timestamptz, + completed_at -> Nullable, + } +} + diesel::table! { providers (id) { id -> Int4, @@ -141,6 +161,10 @@ diesel::joinable!(chat_sessions -> users (user_id)); diesel::joinable!(external_api_tools -> users (user_id)); diesel::joinable!(files -> chat_sessions (session_id)); diesel::joinable!(files -> users (user_id)); +diesel::joinable!(logs -> chat_messages (message_id)); +diesel::joinable!(logs -> chat_sessions (session_id)); +diesel::joinable!(logs -> providers (provider_id)); +diesel::joinable!(logs -> users (user_id)); diesel::joinable!(providers -> secrets (api_key_id)); diesel::joinable!(providers -> users (user_id)); diesel::joinable!(secrets -> users (user_id)); @@ -153,6 +177,7 @@ diesel::allow_tables_to_appear_in_same_query!( chat_sessions, external_api_tools, files, + logs, providers, secrets, system_tools, diff --git a/server-new/src/llm/error.rs b/server-new/src/llm/error.rs index f555cdd..8dde2eb 100644 --- a/server-new/src/llm/error.rs +++ b/server-new/src/llm/error.rs @@ -3,11 +3,11 @@ use crate::services::stream::error::StreamingError; /// Errors that can occur in an LLM provider request #[derive(Debug, thiserror::Error)] pub enum LlmRequestError { - #[error("provider error: {0}")] + #[error("Provider error: {0}")] Provider(String), - #[error("failed to read response: {0}")] + #[error("Failed to read response: {0}")] Read(#[from] reqwest::Error), - #[error("no content")] + #[error("No content")] NoContent, } diff --git a/server-new/src/llm/interface.rs b/server-new/src/llm/interface.rs index 2f9e6ac..7b2f344 100644 --- a/server-new/src/llm/interface.rs +++ b/server-new/src/llm/interface.rs @@ -12,14 +12,33 @@ pub trait LlmProvider: Send + Sync { } /// API response to a prompt request from the LLM provider -pub type LlmPromptResponse<'r> = BoxFuture<'r, Result>; +pub type LlmPromptResponse<'r> = BoxFuture<'r, Result>; /// Initial API response to a streaming request from the LLM provider -pub type LlmStreamingResponse<'r> = BoxFuture<'r, Result>; +pub type LlmStreamingResponse<'r> = + BoxFuture<'r, Result<(LlmStream, LlmResponseMeta), LlmRequestError>>; /// The response stream from the LLM provider pub type LlmStream = BoxStream<'static, LlmStreamChunkResult>; /// The type of the chunks in the LLM response stream pub type LlmStreamChunkResult = Result; +/// Prompt response data from the LLM provider +#[derive(Debug, Default)] +pub struct LlmResponse { + pub text: String, + pub usage: LlmUsage, + pub meta: LlmResponseMeta, +} + +#[derive(Debug, Default)] +pub struct LlmResponseMeta { + pub request_id: Option, +} +impl LlmResponseMeta { + pub fn new(request_id: Option) -> Self { + Self { request_id } + } +} + /// A streaming chunk of data from the LLM provider pub enum LlmStreamChunk { Text(String), diff --git a/server-new/src/llm/providers/anthropic/mod.rs b/server-new/src/llm/providers/anthropic/mod.rs index 2a2dc80..d8bfa9e 100644 --- a/server-new/src/llm/providers/anthropic/mod.rs +++ b/server-new/src/llm/providers/anthropic/mod.rs @@ -4,9 +4,9 @@ use futures::StreamExt; use crate::llm::{ error::LlmRequestError, - interface::{LlmPromptResponse, LlmProvider, LlmStreamingResponse}, + interface::*, providers::utils, - types::{LlmChatRequest, LlmPrompt, LlmUsage}, + types::{LlmChatRequest, LlmPrompt}, }; mod request; @@ -16,6 +16,7 @@ use {request::*, response::*}; const MESSAGES_API_URL: &str = "https://api.anthropic.com/v1/messages"; const API_VERSION: &str = "2023-06-01"; +const REQ_ID_HEADER: &'static str = "request-id"; const DEFAULT_MAX_TOKENS: u32 = 4096; /// Anthropic chat provider @@ -50,18 +51,17 @@ impl LlmProvider for AnthropicProvider { }; Box::pin(async move { - let mut response: AnthropicResponse = utils::llm_api_request( + let raw_response = utils::llm_api_request( self.client .post(MESSAGES_API_URL) .header("anthropic-version", API_VERSION) - .header("content-type", "application/json") .header("x-api-key", &self.api_key) .json(&request), "Anthropic", ) - .await? - .json() .await?; + let request_id = utils::extract_header(&raw_response, REQ_ID_HEADER); + let mut response: AnthropicResponse = raw_response.json().await?; let text = response .content @@ -71,12 +71,12 @@ impl LlmProvider for AnthropicProvider { _ => None, }) .ok_or_else(|| LlmRequestError::NoContent)?; - if let Some(usage) = response.usage { - let usage: LlmUsage = usage.into(); - tracing::info!("Prompt usage: {:?}", usage); - } - Ok(text) + Ok(LlmResponse { + text, + usage: response.usage.map(Into::into).unwrap_or_default(), + meta: LlmResponseMeta::new(request_id), + }) }) } @@ -99,12 +99,12 @@ impl LlmProvider for AnthropicProvider { self.client .post(MESSAGES_API_URL) .header("anthropic-version", API_VERSION) - .header("content-type", "application/json") .header("x-api-key", &self.api_key) .json(&request), "Anthropic", ) .await?; + let request_id = utils::extract_header(&response, REQ_ID_HEADER); let stream = async_stream::stream! { let mut sse_event_stream = utils::get_sse_events(response); @@ -121,7 +121,7 @@ impl LlmProvider for AnthropicProvider { } }; - Ok(stream.boxed()) + Ok((stream.boxed(), LlmResponseMeta::new(request_id))) }) } } diff --git a/server-new/src/llm/providers/lorem.rs b/server-new/src/llm/providers/lorem.rs index 1d1afc3..eb9a05b 100644 --- a/server-new/src/llm/providers/lorem.rs +++ b/server-new/src/llm/providers/lorem.rs @@ -7,11 +7,8 @@ use tokio::time::{Interval, interval}; use crate::llm::{ error::LlmStreamChunkError, - interface::{ - LlmPromptResponse, LlmProvider, LlmStream, LlmStreamChunk, LlmStreamChunkResult, - LlmStreamingResponse, - }, - types::{LlmChatRequest, LlmPrompt}, + interface::*, + types::{LlmChatRequest, LlmPrompt, LlmUsage}, }; /// A test/dummy provider that streams 'lorem ipsum...' and emits test errors during the stream @@ -60,8 +57,18 @@ impl Stream for LoremStream { } impl LlmProvider for LoremProvider { - fn prompt<'r>(&'r self, _prompt: LlmPrompt<'r>) -> LlmPromptResponse<'r> { - Box::pin(async { Ok("Lorem ipsum".to_owned()) }) + fn prompt<'r>(&'r self, prompt: LlmPrompt<'r>) -> LlmPromptResponse<'r> { + let response = LlmResponse { + text: "Lorem ipsum".into(), + usage: LlmUsage { + input_tokens: Some((prompt.text.len() / 4) as u32), + output_tokens: Some(4), + ..Default::default() + }, + ..Default::default() + }; + + Box::pin(async { Ok(response) }) } fn stream_chat<'r>(&'r self, _request: LlmChatRequest<'r>) -> LlmStreamingResponse<'r> { @@ -102,7 +109,7 @@ impl LlmProvider for LoremProvider { }); tokio::time::sleep(Duration::from_millis(1000)).await; // Simulate initial request latency - Ok(stream) + Ok((stream, LlmResponseMeta::default())) }) } } diff --git a/server-new/src/llm/providers/ollama/mod.rs b/server-new/src/llm/providers/ollama/mod.rs index c69ae51..8221741 100644 --- a/server-new/src/llm/providers/ollama/mod.rs +++ b/server-new/src/llm/providers/ollama/mod.rs @@ -4,7 +4,7 @@ use futures::StreamExt; use crate::llm::{ error::LlmRequestError, - interface::{LlmPromptResponse, LlmProvider, LlmStreamingResponse}, + interface::*, providers::utils, types::{LlmChatRequest, LlmPrompt}, }; @@ -51,7 +51,6 @@ impl LlmProvider for OllamaProvider { let res: OllamaCompletionResponse = utils::llm_api_request( self.client .post(format!("{}{}", self.base_url, COMPLETION_API_URL)) - .header("content-type", "application/json") .json(&request), "Ollama", ) @@ -59,14 +58,15 @@ impl LlmProvider for OllamaProvider { .json() .await?; - if let Some(usage) = res.usage() { - tracing::info!("Prompt usage: {:?}", usage); - } if res.response.is_empty() { return Err(LlmRequestError::NoContent); } - Ok(res.response) + Ok(LlmResponse { + usage: res.usage().unwrap_or_default(), + text: res.response, + ..Default::default() + }) }) } @@ -91,7 +91,6 @@ impl LlmProvider for OllamaProvider { let response = utils::llm_api_request( self.client .post(format!("{}{}", self.base_url, CHAT_API_URL)) - .header("content-type", "application/json") .json(&request), "Ollama", ) @@ -120,7 +119,7 @@ impl LlmProvider for OllamaProvider { // } }; - Ok(stream.boxed()) + Ok((stream.boxed(), LlmResponseMeta::default())) }) } } diff --git a/server-new/src/llm/providers/openai/mod.rs b/server-new/src/llm/providers/openai/mod.rs index 298b88c..67dc22f 100644 --- a/server-new/src/llm/providers/openai/mod.rs +++ b/server-new/src/llm/providers/openai/mod.rs @@ -6,9 +6,9 @@ use crate::{ db::models::OpenAISubtype, llm::{ error::LlmRequestError, - interface::{LlmPromptResponse, LlmProvider, LlmStreamingResponse}, + interface::*, providers::utils, - types::{LlmChatRequest, LlmPrompt, LlmUsage}, + types::{LlmChatRequest, LlmPrompt}, }, }; @@ -21,18 +21,24 @@ const OPENAI_API_BASE_URL: &str = "https://api.openai.com/v1"; const OPENROUTER_API_BASE_URL: &str = "https://openrouter.ai/api/v1"; impl OpenAISubtype { - fn name(self) -> &'static str { + fn name(&self) -> &'static str { match self { Self::OpenAI => "OpenAI", Self::OpenRouter => "OpenRouter", } } - fn default_base_url(self) -> &'static str { + fn default_base_url(&self) -> &'static str { match self { Self::OpenAI => OPENAI_API_BASE_URL, Self::OpenRouter => OPENROUTER_API_BASE_URL, } } + fn req_id_header(&self) -> &'static str { + match self { + OpenAISubtype::OpenAI => "X-Request-Id", + OpenAISubtype::OpenRouter => "X-Generation-Id", + } + } fn use_max_completion_tokens(self) -> bool { self == Self::OpenAI } @@ -140,16 +146,16 @@ impl LlmProvider for OpenAIProvider { Box::pin(async move { let provider_name = self.config.subtype.name(); - let mut response: OpenAIResponse = utils::llm_api_request( + let raw_response = utils::llm_api_request( self.client .post(&format!("{}/chat/completions", self.config.base_url)) .bearer_auth(&self.config.api_key) .json(&request), provider_name, ) - .await? - .json() .await?; + let req_id = utils::extract_header(&raw_response, self.config.subtype.req_id_header()); + let mut response: OpenAIResponse = raw_response.json().await?; let text = response .choices @@ -157,12 +163,12 @@ impl LlmProvider for OpenAIProvider { .and_then(|choice| choice.message.as_mut()) .and_then(|message| message.content.take()) .ok_or(LlmRequestError::NoContent)?; - if let Some(usage) = response.usage { - let usage: LlmUsage = usage.into(); - tracing::info!("Prompt usage: {usage:?}"); - } - Ok(text) + Ok(LlmResponse { + text, + usage: response.usage.map(Into::into).unwrap_or_default(), + meta: LlmResponseMeta::new(req_id), + }) }) } @@ -195,6 +201,7 @@ impl LlmProvider for OpenAIProvider { provider_name, ) .await?; + let req_id = utils::extract_header(&response, self.config.subtype.req_id_header()); let stream = async_stream::stream! { let mut sse_event_stream = utils::get_sse_events(response); @@ -220,7 +227,7 @@ impl LlmProvider for OpenAIProvider { // } }; - Ok(stream.boxed()) + Ok((stream.boxed(), LlmResponseMeta::new(req_id))) }) } diff --git a/server-new/src/llm/providers/utils.rs b/server-new/src/llm/providers/utils.rs index 514dbc5..159f9d8 100644 --- a/server-new/src/llm/providers/utils.rs +++ b/server-new/src/llm/providers/utils.rs @@ -75,3 +75,12 @@ pub async fn llm_api_request( Ok(response) } + +/// Convenience function to extract a header (e.g. request ID) from an API response +pub fn extract_header(response: &reqwest::Response, header_name: &str) -> Option { + response + .headers() + .get(header_name) + .and_then(|h| h.to_str().ok()) + .map(str::to_owned) +} diff --git a/server-new/src/llm/types.rs b/server-new/src/llm/types.rs index 8cd2280..40fa6dd 100644 --- a/server-new/src/llm/types.rs +++ b/server-new/src/llm/types.rs @@ -59,7 +59,7 @@ pub struct LlmAssistantMessage { } /// Usage stats from the LLM provider -#[derive(Debug, Default, Serialize, Deserialize, JsonSchema)] +#[derive(Debug, Default, Clone, Copy, Serialize, Deserialize, JsonSchema)] pub struct LlmUsage { pub input_tokens: Option, pub output_tokens: Option, diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index 4309765..1a1cf76 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -7,12 +7,12 @@ use crate::{ db::{ DbPool, DbService, models::{ - AssistantMeta, ChatRsMessage, ChatRsMessageMeta, ChatRsMessageRole, NewChatRsMessage, - UserMeta, + AssistantMeta, ChatRsLogKind, ChatRsLogStatus, ChatRsMessage, ChatRsMessageMeta, + ChatRsMessageRole, NewChatRsMessage, UserMeta, }, }, llm::{ - interface::LlmProvider, + interface::{LlmProvider, LlmResponseMeta}, types::{LlmChatOptions, LlmChatRequest, LlmMessage, LlmUserMessage}, }, services::{ @@ -95,26 +95,64 @@ impl<'r> ChatService<'r> { /// Send a single prompt to the LLM provider and stream the response pub async fn prompt( &self, + db: &mut DbService, user_id: Uuid, + provider_id: i32, provider: Arc, prompt: LlmUserMessage, options: LlmChatOptions, ) -> Result { - let request = LlmChatRequest { - messages: &[LlmMessage::User(prompt)], - options: &options, - }; - let response_stream = provider.stream_chat(request).await?; - let stream_key = StreamingService::prompt_key(&user_id); let (stream_access, ws_writer, ws_reader) = StreamingService::new(self.tinistream) .create_stream(&stream_key) .await?; + let log_id = db + .logs() + .create() + .kind(ChatRsLogKind::Prompt) + .user_id(&user_id) + .provider_id(provider_id) + .model(&options.model) + .build() + .await?; + + let request = LlmChatRequest { + messages: &[LlmMessage::User(prompt)], + options: &options, + }; + let (response_stream, response_meta) = match provider.stream_chat(request).await { + Ok(response) => response, + Err(err) => { + db.logs() + .complete() + .id(log_id) + .status(ChatRsLogStatus::Failed) + .error(&err.to_string()) + .build() + .await?; + return Err(ChatError::Request(err)); + } + }; + // Spawn thread to process LLM streaming response + let db_pool = self.db_pool.to_owned(); let tinistream_client = self.tinistream.to_owned(); tokio::spawn(async move { - let _ = StreamingService::process_stream(response_stream, ws_writer, ws_reader).await; + let output = + StreamingService::process_stream(response_stream, ws_writer, ws_reader).await; + if let Ok(mut db) = DbService::from_pool(&db_pool).await { + let _ = db + .logs() + .complete() + .id(log_id) + .status(output.status()) + .maybe_request_id(response_meta.request_id.as_deref()) + .maybe_usage(output.usage.as_ref()) + .maybe_error(output.errors.map(|e| e.join(", ")).as_deref()) + .build() + .await; + } let _ = StreamingService::new(&tinistream_client) .end_stream(&stream_key) .await; @@ -146,8 +184,9 @@ impl<'r> ChatService<'r> { titles::generate_title( user_id, session_id, - &user_message.text, + provider_id, &provider, + &user_message.text, &chat_options.model, self.db_pool, ); @@ -165,6 +204,7 @@ impl<'r> ChatService<'r> { } self.start_assistant_stream( + db, provider, messages, ChatStreamParams { @@ -201,6 +241,7 @@ impl<'r> ChatService<'r> { .id; self.start_assistant_stream( + db, provider, messages, ChatStreamParams { @@ -217,6 +258,7 @@ impl<'r> ChatService<'r> { /// Start the LLM response stream and return the access URL & token async fn start_assistant_stream( &self, + db: &mut DbService, provider: Arc, messages: Vec, params: ChatStreamParams, @@ -226,14 +268,38 @@ impl<'r> ChatService<'r> { if streams.exists_stream(&stream_key).await? { return Err(ChatError::AlreadyStreaming); } + let (stream_access, ws_writer, ws_reader) = streams.create_stream(&stream_key).await?; - let response_stream = provider + let log_id = db + .logs() + .create() + .kind(ChatRsLogKind::Chat) + .user_id(¶ms.user_id) + .session_id(¶ms.session_id) + .provider_id(params.provider_id) + .model(¶ms.chat_options.model) + .build() + .await?; + + let (response_stream, meta) = match provider .stream_chat(LlmChatRequest { messages: &messages::build_llm_messages(messages)?, options: ¶ms.chat_options, }) - .await?; - let (stream_access, ws_writer, ws_reader) = streams.create_stream(&stream_key).await?; + .await + { + Ok(response) => response, + Err(err) => { + db.logs() + .complete() + .id(log_id) + .status(ChatRsLogStatus::Failed) + .error(&err.to_string()) + .build() + .await?; + return Err(ChatError::Request(err)); + } + }; // Spawn thread to process and save LLM streaming response let db_pool = self.db_pool.to_owned(); @@ -242,7 +308,7 @@ impl<'r> ChatService<'r> { let output = StreamingService::process_stream(response_stream, ws_writer, ws_reader).await; let stream_cancelled = output.cancelled; - if let Err(err) = Self::persist_response(output, params, db_pool).await { + if let Err(err) = Self::persist_response(output, params, log_id, meta, db_pool).await { tracing::error!("Failed to save assistant response: {err}"); } @@ -261,23 +327,26 @@ impl<'r> ChatService<'r> { async fn persist_response( output: LlmStreamOutput, params: ChatStreamParams, + log_id: i32, + meta: LlmResponseMeta, db_pool: DbPool, ) -> Result { let mut db = DbService::from_pool(&db_pool).await?; + let assistant_meta = AssistantMeta { provider_id: params.provider_id, provider_options: Some(params.chat_options), // tool_calls: response.tool_calls, // files: image_ids, usage: output.usage, - errors: output.errors, + errors: output.errors.clone(), partial: output.cancelled.then_some(true), ..Default::default() }; let new_message = db .chats() .save_message(NewChatRsMessage { - content: &output.text.unwrap_or_default(), + content: &output.text.as_deref().unwrap_or_default(), meta: ChatRsMessageMeta::new_assistant(assistant_meta), role: ChatRsMessageRole::Assistant, session_id: ¶ms.session_id, @@ -289,6 +358,17 @@ impl<'r> ChatService<'r> { .await?; } + db.logs() + .complete() + .id(log_id) + .status(output.status()) + .message_id(&new_message.id) + .maybe_request_id(meta.request_id.as_deref()) + .maybe_usage(output.usage.as_ref()) + .maybe_error(output.errors.map(|e| e.join(", ")).as_deref()) + .build() + .await?; + Ok(new_message) } } diff --git a/server-new/src/services/chat/titles.rs b/server-new/src/services/chat/titles.rs index 5b8808f..e6d785b 100644 --- a/server-new/src/services/chat/titles.rs +++ b/server-new/src/services/chat/titles.rs @@ -3,9 +3,12 @@ use std::sync::Arc; use uuid::Uuid; use crate::{ - db::{DbPool, DbService, models::UpdateChatRsSession}, + db::{ + DbPool, DbService, + models::{ChatRsLogKind, ChatRsLogStatus, UpdateChatRsSession}, + }, llm::{ - interface::LlmProvider, + interface::{LlmProvider, LlmResponse}, types::{LlmChatOptions, LlmPrompt}, }, services::chat::error::ChatError, @@ -14,20 +17,22 @@ use crate::{ /// Spawn a task to generate a title for the chat session pub fn generate_title( user_id: Uuid, - session_id: Uuid, - first_message: &str, + sess_id: Uuid, + provider_id: i32, provider: &Arc, + first_message: &str, model: &str, pool: &DbPool, ) { - let user_message = first_message.to_owned(); + let msg = first_message.to_owned(); let provider = Arc::clone(provider); let model = model.to_owned(); let pool = pool.to_owned(); tokio::spawn(async move { - if let Err(err) = generate(user_id, session_id, user_message, provider, model, pool).await { - tracing::warn!("Failed to generate title: {err}"); + if let Err(err) = generate(user_id, sess_id, provider_id, provider, msg, model, pool).await + { + tracing::warn!("Error while generating session title: {err}"); } }); } @@ -40,35 +45,64 @@ const TITLE_PROMPT_MAX_TOKENS: u32 = 20; async fn generate( user_id: Uuid, session_id: Uuid, - user_message: String, + provider_id: i32, provider: Arc, + user_message: String, model: String, db_pool: DbPool, ) -> Result<(), ChatError> { - let prompt = format!("{TITLE_PROMPT}\n\n\"{user_message}\""); - let title = provider - .prompt(LlmPrompt { - text: &prompt, - options: &LlmChatOptions { - model, - temperature: Some(TITLE_PROMPT_TEMPERATURE), - max_tokens: Some(TITLE_PROMPT_MAX_TOKENS), - ..Default::default() - }, - }) - .await?; - let mut db = DbService::from_pool(&db_pool).await?; - db.chats() - .update_session( - &user_id, - &session_id, - UpdateChatRsSession { - title: Some(title.trim()), - ..Default::default() - }, - ) + let log_id = db + .logs() + .create() + .user_id(&user_id) + .session_id(&session_id) + .provider_id(provider_id) + .kind(ChatRsLogKind::Title) + .model(&model) + .build() .await?; + let prompt = LlmPrompt { + text: &format!("{TITLE_PROMPT}\n\n\"{user_message}\""), + options: &LlmChatOptions { + model, + temperature: Some(TITLE_PROMPT_TEMPERATURE), + max_tokens: Some(TITLE_PROMPT_MAX_TOKENS), + ..Default::default() + }, + }; + match provider.prompt(prompt).await { + Ok(LlmResponse { text, usage, meta }) => { + db.logs() + .complete() + .id(log_id) + .status(ChatRsLogStatus::Completed) + .usage(&usage) + .maybe_request_id(meta.request_id.as_deref()) + .build() + .await?; + db.chats() + .update_session( + &user_id, + &session_id, + UpdateChatRsSession { + title: Some(text.trim()), + ..Default::default() + }, + ) + .await?; + } + Err(err) => { + db.logs() + .complete() + .id(log_id) + .status(ChatRsLogStatus::Failed) + .error(&err.to_string()) + .build() + .await?; + } + } + Ok(()) } diff --git a/server-new/src/services/stream/mod.rs b/server-new/src/services/stream/mod.rs index 71b5d7f..0885327 100644 --- a/server-new/src/services/stream/mod.rs +++ b/server-new/src/services/stream/mod.rs @@ -15,6 +15,7 @@ mod writer; mod tests; use crate::{ + db::models::ChatRsLogStatus, llm::{interface::LlmStream, types::LlmUsage}, services::stream::error::StreamingError, }; @@ -33,6 +34,18 @@ pub struct LlmStreamOutput { pub errors: Option>, pub cancelled: bool, } +impl LlmStreamOutput { + /// Get the logged status for this response + pub fn status(&self) -> ChatRsLogStatus { + if self.cancelled { + ChatRsLogStatus::Cancelled + } else if self.errors.is_some() { + ChatRsLogStatus::Failed + } else { + ChatRsLogStatus::Completed + } + } +} type WsWriter = SplitSink; type WsReader = SplitStream; From fd2f654cb0203ef5cd019f08fac2217456b15acf Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 05:13:08 -0400 Subject: [PATCH 087/111] fix regeneration --- server-new/src/services/chat/mod.rs | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index 1a1cf76..fc971ef 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -228,17 +228,16 @@ impl<'r> ChatService<'r> { provider: Arc, chat_options: LlmChatOptions, ) -> Result { - let (chat_session, messages) = db + let (chat_session, mut messages) = db .chats() .find_session_with_messages(&user_id, &session_id) .await?; chat_session.ok_or(ChatError::SessionNotFound)?; - let assistant_message_id = messages - .iter() - .rev() - .find(|m| m.role.is_assistant()) - .ok_or(ChatError::NoAssistantResponse)? - .id; + + let last_message = messages.pop(); + if last_message.as_ref().is_none_or(|m| !m.role.is_assistant()) { + return Err(ChatError::NoAssistantResponse); + } self.start_assistant_stream( db, @@ -249,7 +248,7 @@ impl<'r> ChatService<'r> { session_id, provider_id, chat_options, - replace_message_id: Some(assistant_message_id), + replace_message_id: last_message.map(|m| m.id), }, ) .await From e3f441798da0c5b1343ef8237e2c377cf4eb69fb Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 05:17:52 -0400 Subject: [PATCH 088/111] fix order of operations in chat streams --- server-new/src/services/chat/mod.rs | 11 +++++------ 1 file changed, 5 insertions(+), 6 deletions(-) diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index fc971ef..db47504 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -102,11 +102,6 @@ impl<'r> ChatService<'r> { prompt: LlmUserMessage, options: LlmChatOptions, ) -> Result { - let stream_key = StreamingService::prompt_key(&user_id); - let (stream_access, ws_writer, ws_reader) = StreamingService::new(self.tinistream) - .create_stream(&stream_key) - .await?; - let log_id = db .logs() .create() @@ -136,6 +131,10 @@ impl<'r> ChatService<'r> { }; // Spawn thread to process LLM streaming response + let stream_key = StreamingService::prompt_key(&user_id); + let (stream_access, ws_writer, ws_reader) = StreamingService::new(self.tinistream) + .create_stream(&stream_key) + .await?; let db_pool = self.db_pool.to_owned(); let tinistream_client = self.tinistream.to_owned(); tokio::spawn(async move { @@ -267,7 +266,6 @@ impl<'r> ChatService<'r> { if streams.exists_stream(&stream_key).await? { return Err(ChatError::AlreadyStreaming); } - let (stream_access, ws_writer, ws_reader) = streams.create_stream(&stream_key).await?; let log_id = db .logs() @@ -301,6 +299,7 @@ impl<'r> ChatService<'r> { }; // Spawn thread to process and save LLM streaming response + let (stream_access, ws_writer, ws_reader) = streams.create_stream(&stream_key).await?; let db_pool = self.db_pool.to_owned(); let tinistream_client = self.tinistream.to_owned(); tokio::spawn(async move { From b28edaf4f548df55cac6e9f833010509635b2ec5 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 05:20:17 -0400 Subject: [PATCH 089/111] rename llm_logs table --- .../2026-07-15-053539-0000_add_request_logs/up.sql | 10 +++++----- server-new/src/db/models/log.rs | 6 +++--- server-new/src/db/repositories/log.rs | 12 ++++++------ server-new/src/db/schema.rs | 12 ++++++------ 4 files changed, 20 insertions(+), 20 deletions(-) diff --git a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql index c5cf98e..ca168f9 100644 --- a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql +++ b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql @@ -1,4 +1,4 @@ -CREATE TABLE logs ( +CREATE TABLE llm_logs ( id SERIAL PRIMARY KEY, kind text NOT NULL, -- chat, title, prompt, image, audio, etc. user_id uuid NOT NULL REFERENCES users (id) ON DELETE CASCADE, @@ -16,10 +16,10 @@ CREATE TABLE logs ( completed_at timestamptz ); -CREATE INDEX logs_user_id_started_at_idx ON logs (user_id, started_at DESC); +CREATE INDEX llm_logs_user_id_started_at_idx ON llm_logs (user_id, started_at DESC); -CREATE INDEX logs_session_id_idx ON logs (session_id); +CREATE INDEX llm_logs_session_id_idx ON llm_logs (session_id); -CREATE INDEX logs_message_id_idx ON logs (message_id); +CREATE INDEX llm_logs_message_id_idx ON llm_logs (message_id); -CREATE INDEX logs_provider_id_started_at_idx ON logs (provider_id, started_at DESC); +CREATE INDEX llm_logs_provider_id_started_at_idx ON llm_logs (provider_id, started_at DESC); diff --git a/server-new/src/db/models/log.rs b/server-new/src/db/models/log.rs index c2635fa..eecbda6 100644 --- a/server-new/src/db/models/log.rs +++ b/server-new/src/db/models/log.rs @@ -9,7 +9,7 @@ use crate::db::{UtcDateTime, models::ChatRsUser}; #[derive(Identifiable, Associations, Queryable, Selectable, Serialize, JsonSchema)] #[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] -#[diesel(table_name = super::schema::logs)] +#[diesel(table_name = super::schema::llm_logs)] pub struct ChatRsLog { pub id: i32, pub kind: String, @@ -47,7 +47,7 @@ pub enum ChatRsLogStatus { } #[derive(Insertable)] -#[diesel(table_name = super::schema::logs)] +#[diesel(table_name = super::schema::llm_logs)] pub struct NewChatRsLog<'a> { pub kind: &'a str, pub user_id: &'a Uuid, @@ -58,7 +58,7 @@ pub struct NewChatRsLog<'a> { } #[derive(Default, AsChangeset)] -#[diesel(table_name = super::schema::logs)] +#[diesel(table_name = super::schema::llm_logs)] pub struct UpdateChatRsLog<'a> { pub message_id: Option<&'a Uuid>, pub request_id: Option<&'a str>, diff --git a/server-new/src/db/repositories/log.rs b/server-new/src/db/repositories/log.rs index d53d869..44021d7 100644 --- a/server-new/src/db/repositories/log.rs +++ b/server-new/src/db/repositories/log.rs @@ -9,7 +9,7 @@ use crate::{ db::{ DbConnection, models::{ChatRsLogKind, ChatRsLogStatus, NewChatRsLog, UpdateChatRsLog}, - schema::logs, + schema::llm_logs, }, llm::types::LlmUsage, }; @@ -43,9 +43,9 @@ impl<'a> LogRepository<'a> { status: ChatRsLogStatus::Started.as_ref(), }; - diesel::insert_into(logs::table) + diesel::insert_into(llm_logs::table) .values(new_log) - .returning(logs::id) + .returning(llm_logs::id) .get_result(self.db) .await } @@ -76,10 +76,10 @@ impl<'a> LogRepository<'a> { completed_at: Utc::now(), }; - diesel::update(logs::table) - .filter(logs::id.eq(id)) + diesel::update(llm_logs::table) + .filter(llm_logs::id.eq(id)) .set(update_log) - .returning(logs::id) + .returning(llm_logs::id) .get_result(self.db) .await } diff --git a/server-new/src/db/schema.rs b/server-new/src/db/schema.rs index 9e6218b..4858351 100644 --- a/server-new/src/db/schema.rs +++ b/server-new/src/db/schema.rs @@ -85,7 +85,7 @@ diesel::table! { } diesel::table! { - logs (id) { + llm_logs (id) { id -> Int4, kind -> Text, user_id -> Uuid, @@ -161,10 +161,10 @@ diesel::joinable!(chat_sessions -> users (user_id)); diesel::joinable!(external_api_tools -> users (user_id)); diesel::joinable!(files -> chat_sessions (session_id)); diesel::joinable!(files -> users (user_id)); -diesel::joinable!(logs -> chat_messages (message_id)); -diesel::joinable!(logs -> chat_sessions (session_id)); -diesel::joinable!(logs -> providers (provider_id)); -diesel::joinable!(logs -> users (user_id)); +diesel::joinable!(llm_logs -> chat_messages (message_id)); +diesel::joinable!(llm_logs -> chat_sessions (session_id)); +diesel::joinable!(llm_logs -> providers (provider_id)); +diesel::joinable!(llm_logs -> users (user_id)); diesel::joinable!(providers -> secrets (api_key_id)); diesel::joinable!(providers -> users (user_id)); diesel::joinable!(secrets -> users (user_id)); @@ -177,7 +177,7 @@ diesel::allow_tables_to_appear_in_same_query!( chat_sessions, external_api_tools, files, - logs, + llm_logs, providers, secrets, system_tools, From d1ceae7d2680026b402ccde26069d12d59d90749 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 05:30:29 -0400 Subject: [PATCH 090/111] accurate completion time for logs --- server-new/src/db/repositories/log.rs | 6 +++--- server-new/src/services/chat/mod.rs | 2 ++ server-new/src/services/chat/titles.rs | 18 ++++++++++-------- 3 files changed, 15 insertions(+), 11 deletions(-) diff --git a/server-new/src/db/repositories/log.rs b/server-new/src/db/repositories/log.rs index 44021d7..8437e53 100644 --- a/server-new/src/db/repositories/log.rs +++ b/server-new/src/db/repositories/log.rs @@ -1,13 +1,12 @@ use bigdecimal::{BigDecimal, FromPrimitive}; use bon::bon; -use chrono::Utc; use diesel::prelude::*; use diesel_async::RunQueryDsl; use uuid::Uuid; use crate::{ db::{ - DbConnection, + DbConnection, UtcDateTime, models::{ChatRsLogKind, ChatRsLogStatus, NewChatRsLog, UpdateChatRsLog}, schema::llm_logs, }, @@ -60,6 +59,7 @@ impl<'a> LogRepository<'a> { usage: Option<&LlmUsage>, error: Option<&str>, status: ChatRsLogStatus, + completed_at: Option, ) -> QueryResult { let update_log = UpdateChatRsLog { message_id, @@ -73,7 +73,7 @@ impl<'a> LogRepository<'a> { cost: usage.and_then(|u| u.cost.and_then(BigDecimal::from_f32)), error, status: status.as_ref(), - completed_at: Utc::now(), + completed_at: completed_at.unwrap_or_else(|| chrono::Utc::now()), }; diesel::update(llm_logs::table) diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index db47504..364c886 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -329,6 +329,7 @@ impl<'r> ChatService<'r> { meta: LlmResponseMeta, db_pool: DbPool, ) -> Result { + let completed_at = chrono::Utc::now(); let mut db = DbService::from_pool(&db_pool).await?; let assistant_meta = AssistantMeta { @@ -360,6 +361,7 @@ impl<'r> ChatService<'r> { .complete() .id(log_id) .status(output.status()) + .completed_at(completed_at) .message_id(&new_message.id) .maybe_request_id(meta.request_id.as_deref()) .maybe_usage(output.usage.as_ref()) diff --git a/server-new/src/services/chat/titles.rs b/server-new/src/services/chat/titles.rs index e6d785b..1700941 100644 --- a/server-new/src/services/chat/titles.rs +++ b/server-new/src/services/chat/titles.rs @@ -74,14 +74,7 @@ async fn generate( }; match provider.prompt(prompt).await { Ok(LlmResponse { text, usage, meta }) => { - db.logs() - .complete() - .id(log_id) - .status(ChatRsLogStatus::Completed) - .usage(&usage) - .maybe_request_id(meta.request_id.as_deref()) - .build() - .await?; + let completed_at = chrono::Utc::now(); db.chats() .update_session( &user_id, @@ -92,6 +85,15 @@ async fn generate( }, ) .await?; + db.logs() + .complete() + .id(log_id) + .status(ChatRsLogStatus::Completed) + .completed_at(completed_at) + .usage(&usage) + .maybe_request_id(meta.request_id.as_deref()) + .build() + .await?; } Err(err) => { db.logs() From efeefbb924fe8ffa44f0baea19ca71159431a448 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 05:31:20 -0400 Subject: [PATCH 091/111] Update down.sql --- .../migrations/2026-07-15-053539-0000_add_request_logs/down.sql | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/down.sql b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/down.sql index f2aff2c..e30a71a 100644 --- a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/down.sql +++ b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/down.sql @@ -1 +1 @@ -DROP TABLE logs; +DROP TABLE llm_logs; From ed4bf105dcde19cb80b7df5eade26c557bb093c8 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 06:05:58 -0400 Subject: [PATCH 092/111] lint: address clippy warnings --- server-new/src/db/repositories/log.rs | 2 +- server-new/src/extractors/user.rs | 6 +++--- server-new/src/llm/providers/anthropic/mod.rs | 4 ++-- server-new/src/llm/providers/anthropic/response.rs | 6 ++---- server-new/src/llm/providers/lorem.rs | 2 +- server-new/src/llm/providers/ollama/mod.rs | 3 +-- server-new/src/llm/providers/openai/mod.rs | 6 +++--- server-new/src/services/auth/api_key.rs | 5 +---- server-new/src/services/auth/encryption.rs | 2 +- server-new/src/services/auth/oauth.rs | 4 ++-- server-new/src/services/auth/oauth/discord.rs | 2 +- server-new/src/services/auth/oauth/github.rs | 2 +- server-new/src/services/auth/oauth/google.rs | 2 +- server-new/src/services/auth/oauth/oidc.rs | 2 +- server-new/src/services/auth/proxy.rs | 11 ++++++----- server-new/src/services/auth/session.rs | 2 +- server-new/src/services/auth/session_store.rs | 2 +- server-new/src/services/chat/mod.rs | 2 +- server-new/src/services/chat/titles.rs | 1 - server-new/src/services/model/mod.rs | 3 +-- server-new/src/services/provider/mod.rs | 6 +++--- server-new/src/services/stream/mod.rs | 4 ++-- 22 files changed, 36 insertions(+), 43 deletions(-) diff --git a/server-new/src/db/repositories/log.rs b/server-new/src/db/repositories/log.rs index 8437e53..36ae70f 100644 --- a/server-new/src/db/repositories/log.rs +++ b/server-new/src/db/repositories/log.rs @@ -73,7 +73,7 @@ impl<'a> LogRepository<'a> { cost: usage.and_then(|u| u.cost.and_then(BigDecimal::from_f32)), error, status: status.as_ref(), - completed_at: completed_at.unwrap_or_else(|| chrono::Utc::now()), + completed_at: completed_at.unwrap_or_else(chrono::Utc::now), }; diesel::update(llm_logs::table) diff --git a/server-new/src/extractors/user.rs b/server-new/src/extractors/user.rs index 2b47ef1..977a09d 100644 --- a/server-new/src/extractors/user.rs +++ b/server-new/src/extractors/user.rs @@ -13,9 +13,9 @@ use crate::{db::DbService, error::AppError, state::AppState}; Represents an active user, extracted from the session, proxy headers, or API key. This can be used as an extractor in route handlers: - If used as `CurrentUser`, request will automatically return an unauthorized error -if there is no active user. + if there is no active user. - If used as `Option`, will be `Some` if there is an active user -and `None` otherwise. + and `None` otherwise. */ #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CurrentUser { @@ -72,7 +72,7 @@ impl OptionalFromRequestParts for CurrentUser { .extensions .get::() .ok_or_else(|| AppError::internal(anyhow!("session not attached to request")))?; - let maybe_user_id = auth_service.session().active_user_id(&session).await?; + let maybe_user_id = auth_service.session().active_user_id(session).await?; Ok(maybe_user_id.map(Self::new)) } diff --git a/server-new/src/llm/providers/anthropic/mod.rs b/server-new/src/llm/providers/anthropic/mod.rs index d8bfa9e..deeeab4 100644 --- a/server-new/src/llm/providers/anthropic/mod.rs +++ b/server-new/src/llm/providers/anthropic/mod.rs @@ -16,7 +16,7 @@ use {request::*, response::*}; const MESSAGES_API_URL: &str = "https://api.anthropic.com/v1/messages"; const API_VERSION: &str = "2023-06-01"; -const REQ_ID_HEADER: &'static str = "request-id"; +const REQ_ID_HEADER: &str = "request-id"; const DEFAULT_MAX_TOKENS: u32 = 4096; /// Anthropic chat provider @@ -81,7 +81,7 @@ impl LlmProvider for AnthropicProvider { } fn stream_chat<'r>(&'r self, req: LlmChatRequest<'r>) -> LlmStreamingResponse<'r> { - let (anthropic_messages, system_prompt) = build_anthropic_messages(&req.messages); + let (anthropic_messages, system_prompt) = build_anthropic_messages(req.messages); // let anthropic_tools = tools.as_ref().map(|t| build_anthropic_tools(t)); let request = AnthropicRequest { model: &req.options.model, diff --git a/server-new/src/llm/providers/anthropic/response.rs b/server-new/src/llm/providers/anthropic/response.rs index b71568e..804acbd 100644 --- a/server-new/src/llm/providers/anthropic/response.rs +++ b/server-new/src/llm/providers/anthropic/response.rs @@ -60,10 +60,8 @@ pub fn parse_anthropic_event( // } // } } - AnthropicStreamEvent::MessageDelta { usage } => { - if let Some(usage) = usage { - return Some(Ok(LlmStreamChunk::Usage(usage.into()))); - } + AnthropicStreamEvent::MessageDelta { usage: Some(usage) } => { + return Some(Ok(LlmStreamChunk::Usage(usage.into()))); } AnthropicStreamEvent::Error { error } => { let error_msg = format!("{}: {}", error.error_type, error.message); diff --git a/server-new/src/llm/providers/lorem.rs b/server-new/src/llm/providers/lorem.rs index eb9a05b..a1cc1a9 100644 --- a/server-new/src/llm/providers/lorem.rs +++ b/server-new/src/llm/providers/lorem.rs @@ -43,7 +43,7 @@ impl Stream for LoremStream { std::task::Poll::Ready(_) => { let word = self.words[self.index]; self.index += 1; - if self.index == 0 || self.index % 10 != 0 { + if self.index == 0 || !self.index.is_multiple_of(10) { std::task::Poll::Ready(Some(Ok(LlmStreamChunk::Text(word.to_owned())))) } else { std::task::Poll::Ready(Some(Err(LlmStreamChunkError::Provider( diff --git a/server-new/src/llm/providers/ollama/mod.rs b/server-new/src/llm/providers/ollama/mod.rs index 8221741..232c846 100644 --- a/server-new/src/llm/providers/ollama/mod.rs +++ b/server-new/src/llm/providers/ollama/mod.rs @@ -71,7 +71,7 @@ impl LlmProvider for OllamaProvider { } fn stream_chat<'r>(&'r self, req: LlmChatRequest<'r>) -> LlmStreamingResponse<'r> { - let ollama_messages = build_ollama_messages(&req.messages); + let ollama_messages = build_ollama_messages(req.messages); // let ollama_tools = tools.as_ref().map(|t| build_ollama_tools(t)); let ollama_options = OllamaOptions { temperature: req.options.temperature, @@ -84,7 +84,6 @@ impl LlmProvider for OllamaProvider { // tools: ollama_tools, stream: Some(true), options: Some(ollama_options), - ..Default::default() }; Box::pin(async move { diff --git a/server-new/src/llm/providers/openai/mod.rs b/server-new/src/llm/providers/openai/mod.rs index 67dc22f..68ff002 100644 --- a/server-new/src/llm/providers/openai/mod.rs +++ b/server-new/src/llm/providers/openai/mod.rs @@ -148,7 +148,7 @@ impl LlmProvider for OpenAIProvider { let provider_name = self.config.subtype.name(); let raw_response = utils::llm_api_request( self.client - .post(&format!("{}/chat/completions", self.config.base_url)) + .post(format!("{}/chat/completions", self.config.base_url)) .bearer_auth(&self.config.api_key) .json(&request), provider_name, @@ -174,7 +174,7 @@ impl LlmProvider for OpenAIProvider { fn stream_chat<'r>(&'r self, req: LlmChatRequest<'r>) -> LlmStreamingResponse<'r> { let policy = OpenAIRequestPolicy::new(self.config.subtype); - let openai_messages = build_openai_messages(&req.messages); + let openai_messages = build_openai_messages(req.messages); // let openai_tools = tools.as_ref().map(|t| build_openai_tools(t)); // let request = OpenAIRequest { @@ -195,7 +195,7 @@ impl LlmProvider for OpenAIProvider { Box::pin(async move { let response = utils::llm_api_request( self.client - .post(&format!("{}/chat/completions", self.config.base_url)) + .post(format!("{}/chat/completions", self.config.base_url)) .bearer_auth(&self.config.api_key) .json(&request), provider_name, diff --git a/server-new/src/services/auth/api_key.rs b/server-new/src/services/auth/api_key.rs index 88cda9f..a8c9833 100644 --- a/server-new/src/services/auth/api_key.rs +++ b/server-new/src/services/auth/api_key.rs @@ -38,10 +38,7 @@ impl<'r> ApiKeyService<'r> { ) -> AuthResult<(Uuid, String)> { let key_id = db .api_keys() - .create(NewChatRsApiKey { - user_id: &user_id, - name: &name, - }) + .create(NewChatRsApiKey { user_id, name }) .await?; let (ciphertext, nonce) = self.encryptor.encrypt_bytes(key_id.as_bytes())?; diff --git a/server-new/src/services/auth/encryption.rs b/server-new/src/services/auth/encryption.rs index 0163445..19b4256 100644 --- a/server-new/src/services/auth/encryption.rs +++ b/server-new/src/services/auth/encryption.rs @@ -60,7 +60,7 @@ impl Encryptor { .decrypt(&nonce, ciphertext) .map_err(|_| EncryptorError::Decryption)?; - Ok(String::from_utf8(plaintext).map_err(|_| EncryptorError::Decryption)?) + String::from_utf8(plaintext).map_err(|_| EncryptorError::Decryption) } /// Decrypts a byte slice using AES-256-GCM. diff --git a/server-new/src/services/auth/oauth.rs b/server-new/src/services/auth/oauth.rs index 8a157ce..bf1f2ca 100644 --- a/server-new/src/services/auth/oauth.rs +++ b/server-new/src/services/auth/oauth.rs @@ -80,12 +80,12 @@ impl<'a> OAuthService<'a> { fn oauth_provider( &self, provider: &OAuthProviderEnum, - ) -> AuthResult<(&OAuthClient, &Box)> { + ) -> AuthResult<(&OAuthClient, &dyn OAuthProvider)> { let (client, provider) = self .provider_map .get(provider) .ok_or_else(|| AuthError::BadRequest("unsupported OAuth provider"))?; - Ok((client, provider)) + Ok((client, provider.as_ref())) } fn get_redirect_url(&self, callback_path: &str) -> String { diff --git a/server-new/src/services/auth/oauth/discord.rs b/server-new/src/services/auth/oauth/discord.rs index d6196df..df44f74 100644 --- a/server-new/src/services/auth/oauth/discord.rs +++ b/server-new/src/services/auth/oauth/discord.rs @@ -62,7 +62,7 @@ impl OAuthProvider for DiscordProvider { fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { discord_id: Some(&user_data.id), - name: &user_data + name: user_data .name .as_deref() .or(user_data.username.as_deref()) diff --git a/server-new/src/services/auth/oauth/github.rs b/server-new/src/services/auth/oauth/github.rs index 35314a7..e4885f2 100644 --- a/server-new/src/services/auth/oauth/github.rs +++ b/server-new/src/services/auth/oauth/github.rs @@ -62,7 +62,7 @@ impl OAuthProvider for GitHubProvider { fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { github_id: Some(&user_data.id), - name: &user_data + name: user_data .name .as_deref() .or(user_data.username.as_deref()) diff --git a/server-new/src/services/auth/oauth/google.rs b/server-new/src/services/auth/oauth/google.rs index a81c8bb..8203c79 100644 --- a/server-new/src/services/auth/oauth/google.rs +++ b/server-new/src/services/auth/oauth/google.rs @@ -59,7 +59,7 @@ impl OAuthProvider for GoogleProvider { fn create_new_user<'a>(&self, user_data: &'a UserInfo) -> NewChatRsUser<'a> { NewChatRsUser { google_id: Some(&user_data.id), - name: &user_data + name: user_data .name .as_deref() .or(user_data.username.as_deref()) diff --git a/server-new/src/services/auth/oauth/oidc.rs b/server-new/src/services/auth/oauth/oidc.rs index efbfec0..065eb74 100644 --- a/server-new/src/services/auth/oauth/oidc.rs +++ b/server-new/src/services/auth/oauth/oidc.rs @@ -79,7 +79,7 @@ impl OAuthProvider for OidcProvider { ) -> crate::db::models::NewChatRsUser<'a> { NewChatRsUser { google_id: Some(&user_info.id), - name: &user_info + name: user_info .name .as_deref() .or(user_info.username.as_deref()) diff --git a/server-new/src/services/auth/proxy.rs b/server-new/src/services/auth/proxy.rs index c358ba3..52335ef 100644 --- a/server-new/src/services/auth/proxy.rs +++ b/server-new/src/services/auth/proxy.rs @@ -64,10 +64,10 @@ impl<'r> ProxyService<'r> { .and_then(|groups| groups.to_str().ok()) .unwrap_or_default(); - if let Some(ref allowed_groups) = self.config.user_groups { - if !is_proxy_user_allowed(groups, allowed_groups) { - return Err(AuthError::Unauthorized("proxy user not in allowed group")); - } + if let Some(ref allowed_groups) = self.config.user_groups + && !is_proxy_user_allowed(groups, allowed_groups) + { + return Err(AuthError::Unauthorized("proxy user not in allowed group")); } Ok(Some(ProxyUser { @@ -111,5 +111,6 @@ fn is_proxy_user_allowed(user_groups: &str, allowed_groups: &[String]) -> bool { return true; } } - return false; + + false } diff --git a/server-new/src/services/auth/session.rs b/server-new/src/services/auth/session.rs index 039500b..5f469f3 100644 --- a/server-new/src/services/auth/session.rs +++ b/server-new/src/services/auth/session.rs @@ -69,7 +69,7 @@ impl AuthSessionService { // Cleanup expired sessions #[tracing::instrument(skip(db_pool), level = "debug")] pub async fn session_cleanup(db_pool: &DbPool) -> AuthResult { - let mut db = DbService::from_pool(&db_pool).await?; + let mut db = DbService::from_pool(db_pool).await?; Ok(db.auth_sessions().delete_expired().await?) } } diff --git a/server-new/src/services/auth/session_store.rs b/server-new/src/services/auth/session_store.rs index 24e97d8..52afb61 100644 --- a/server-new/src/services/auth/session_store.rs +++ b/server-new/src/services/auth/session_store.rs @@ -83,7 +83,7 @@ impl SessionStore for SessionDbStore { /// does not exist or has been invalidated (e.g., expired), `None` is /// returned. async fn load(&self, session_id: &Id) -> Result> { - let session_id = Self::get_session_uuid(&session_id); + let session_id = Self::get_session_uuid(session_id); let mut db = self.get_db().await?; match db.auth_sessions().find_active_by_id(&session_id).await { diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index 364c886..c0728ca 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -345,7 +345,7 @@ impl<'r> ChatService<'r> { let new_message = db .chats() .save_message(NewChatRsMessage { - content: &output.text.as_deref().unwrap_or_default(), + content: output.text.as_deref().unwrap_or_default(), meta: ChatRsMessageMeta::new_assistant(assistant_meta), role: ChatRsMessageRole::Assistant, session_id: ¶ms.session_id, diff --git a/server-new/src/services/chat/titles.rs b/server-new/src/services/chat/titles.rs index 1700941..1021a6f 100644 --- a/server-new/src/services/chat/titles.rs +++ b/server-new/src/services/chat/titles.rs @@ -69,7 +69,6 @@ async fn generate( model, temperature: Some(TITLE_PROMPT_TEMPERATURE), max_tokens: Some(TITLE_PROMPT_MAX_TOKENS), - ..Default::default() }, }; match provider.prompt(prompt).await { diff --git a/server-new/src/services/model/mod.rs b/server-new/src/services/model/mod.rs index d746045..e15a3c1 100644 --- a/server-new/src/services/model/mod.rs +++ b/server-new/src/services/model/mod.rs @@ -86,8 +86,7 @@ impl<'r> ModelService<'r> { .remove(provider.as_ref()) .ok_or_else(|| ModelError::ModelsDevProviderNotFound(provider.into()))? .models - .into_iter() - .map(|(_, model)| model) + .into_values() .collect(); let provider_models_str = serde_json::to_string(&provider_models)?; cache.insert(provider.as_ref().to_owned(), provider_models_str); diff --git a/server-new/src/services/provider/mod.rs b/server-new/src/services/provider/mod.rs index 1082108..5e42810 100644 --- a/server-new/src/services/provider/mod.rs +++ b/server-new/src/services/provider/mod.rs @@ -177,7 +177,7 @@ impl<'r> ProviderService<'r> { }; let updated = db .providers() - .update(&user_id, provider_id, update_provider) + .update(user_id, provider_id, update_provider) .await?; Ok(updated) @@ -191,9 +191,9 @@ impl<'r> ProviderService<'r> { ) -> Result { let (_provider, _, api_key_secret) = self.get_provider(db, user_id, provider_id).await?; if let Some(secret) = api_key_secret { - db.secrets().delete(&user_id, &secret.id).await?; + db.secrets().delete(user_id, &secret.id).await?; } - let deleted = db.providers().delete(&user_id, provider_id).await?; + let deleted = db.providers().delete(user_id, provider_id).await?; Ok(deleted) } diff --git a/server-new/src/services/stream/mod.rs b/server-new/src/services/stream/mod.rs index 0885327..be00087 100644 --- a/server-new/src/services/stream/mod.rs +++ b/server-new/src/services/stream/mod.rs @@ -72,13 +72,13 @@ impl<'r> StreamingService<'r> { /// Extract the session ID from the user's stream key pub fn session_id_from_stream_key(key: &str, key_prefix: &str) -> Option { - key.strip_prefix(&key_prefix) + key.strip_prefix(key_prefix) .and_then(|session_id| Uuid::try_parse(session_id).ok()) } /// Check for existing active client stream pub async fn exists_stream(&self, stream_key: &str) -> Result { - Ok(self.tinistream.stream_exists(&stream_key).await?) + Ok(self.tinistream.stream_exists(stream_key).await?) } /// Currently active streams with the given prefix From 3c3adb737c3f6226f7c783817cb39abd9569e15a Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 06:14:46 -0400 Subject: [PATCH 093/111] renaming log status for clarity --- server-new/src/db/models/log.rs | 2 +- server-new/src/llm/error.rs | 2 +- server-new/src/services/chat/mod.rs | 4 ++-- server-new/src/services/chat/titles.rs | 2 +- server-new/src/services/stream/mod.rs | 2 +- 5 files changed, 6 insertions(+), 6 deletions(-) diff --git a/server-new/src/db/models/log.rs b/server-new/src/db/models/log.rs index eecbda6..914f2b3 100644 --- a/server-new/src/db/models/log.rs +++ b/server-new/src/db/models/log.rs @@ -42,8 +42,8 @@ pub enum ChatRsLogKind { pub enum ChatRsLogStatus { Started, Completed, - Failed, Cancelled, + Error, } #[derive(Insertable)] diff --git a/server-new/src/llm/error.rs b/server-new/src/llm/error.rs index 8dde2eb..f161471 100644 --- a/server-new/src/llm/error.rs +++ b/server-new/src/llm/error.rs @@ -3,7 +3,7 @@ use crate::services::stream::error::StreamingError; /// Errors that can occur in an LLM provider request #[derive(Debug, thiserror::Error)] pub enum LlmRequestError { - #[error("Provider error: {0}")] + #[error("{0}")] Provider(String), #[error("Failed to read response: {0}")] Read(#[from] reqwest::Error), diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index c0728ca..58c949f 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -122,7 +122,7 @@ impl<'r> ChatService<'r> { db.logs() .complete() .id(log_id) - .status(ChatRsLogStatus::Failed) + .status(ChatRsLogStatus::Error) .error(&err.to_string()) .build() .await?; @@ -290,7 +290,7 @@ impl<'r> ChatService<'r> { db.logs() .complete() .id(log_id) - .status(ChatRsLogStatus::Failed) + .status(ChatRsLogStatus::Error) .error(&err.to_string()) .build() .await?; diff --git a/server-new/src/services/chat/titles.rs b/server-new/src/services/chat/titles.rs index 1021a6f..cd8e5ff 100644 --- a/server-new/src/services/chat/titles.rs +++ b/server-new/src/services/chat/titles.rs @@ -98,7 +98,7 @@ async fn generate( db.logs() .complete() .id(log_id) - .status(ChatRsLogStatus::Failed) + .status(ChatRsLogStatus::Error) .error(&err.to_string()) .build() .await?; diff --git a/server-new/src/services/stream/mod.rs b/server-new/src/services/stream/mod.rs index be00087..06da595 100644 --- a/server-new/src/services/stream/mod.rs +++ b/server-new/src/services/stream/mod.rs @@ -40,7 +40,7 @@ impl LlmStreamOutput { if self.cancelled { ChatRsLogStatus::Cancelled } else if self.errors.is_some() { - ChatRsLogStatus::Failed + ChatRsLogStatus::Error } else { ChatRsLogStatus::Completed } From 851b3d7b69538f5c3d08d3ce4c40f6118e0a1c1c Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 06:53:24 -0400 Subject: [PATCH 094/111] save request ID from provider request errors --- server-new/src/llm/error.rs | 11 ++++++++++- server-new/src/llm/providers/anthropic/mod.rs | 2 ++ server-new/src/llm/providers/ollama/mod.rs | 2 ++ server-new/src/llm/providers/openai/mod.rs | 2 ++ server-new/src/llm/providers/utils.rs | 16 +++++++++------- server-new/src/services/chat/mod.rs | 2 ++ server-new/src/services/chat/titles.rs | 1 + 7 files changed, 28 insertions(+), 8 deletions(-) diff --git a/server-new/src/llm/error.rs b/server-new/src/llm/error.rs index f161471..eb98e88 100644 --- a/server-new/src/llm/error.rs +++ b/server-new/src/llm/error.rs @@ -3,13 +3,22 @@ use crate::services::stream::error::StreamingError; /// Errors that can occur in an LLM provider request #[derive(Debug, thiserror::Error)] pub enum LlmRequestError { + /// Provider error message with optional request ID #[error("{0}")] - Provider(String), + Provider(String, Option), #[error("Failed to read response: {0}")] Read(#[from] reqwest::Error), #[error("No content")] NoContent, } +impl LlmRequestError { + pub fn req_id(&self) -> Option<&str> { + match self { + LlmRequestError::Provider(_, req_id) => req_id.as_deref(), + _ => None, + } + } +} /// Errors that can occur in an LLM stream chunk #[derive(Debug, thiserror::Error)] diff --git a/server-new/src/llm/providers/anthropic/mod.rs b/server-new/src/llm/providers/anthropic/mod.rs index deeeab4..864ab40 100644 --- a/server-new/src/llm/providers/anthropic/mod.rs +++ b/server-new/src/llm/providers/anthropic/mod.rs @@ -58,6 +58,7 @@ impl LlmProvider for AnthropicProvider { .header("x-api-key", &self.api_key) .json(&request), "Anthropic", + Some(REQ_ID_HEADER), ) .await?; let request_id = utils::extract_header(&raw_response, REQ_ID_HEADER); @@ -102,6 +103,7 @@ impl LlmProvider for AnthropicProvider { .header("x-api-key", &self.api_key) .json(&request), "Anthropic", + Some(REQ_ID_HEADER), ) .await?; let request_id = utils::extract_header(&response, REQ_ID_HEADER); diff --git a/server-new/src/llm/providers/ollama/mod.rs b/server-new/src/llm/providers/ollama/mod.rs index 232c846..a38be20 100644 --- a/server-new/src/llm/providers/ollama/mod.rs +++ b/server-new/src/llm/providers/ollama/mod.rs @@ -53,6 +53,7 @@ impl LlmProvider for OllamaProvider { .post(format!("{}{}", self.base_url, COMPLETION_API_URL)) .json(&request), "Ollama", + None, ) .await? .json() @@ -92,6 +93,7 @@ impl LlmProvider for OllamaProvider { .post(format!("{}{}", self.base_url, CHAT_API_URL)) .json(&request), "Ollama", + None, ) .await?; let stream = async_stream::stream! { diff --git a/server-new/src/llm/providers/openai/mod.rs b/server-new/src/llm/providers/openai/mod.rs index 68ff002..a22af2d 100644 --- a/server-new/src/llm/providers/openai/mod.rs +++ b/server-new/src/llm/providers/openai/mod.rs @@ -152,6 +152,7 @@ impl LlmProvider for OpenAIProvider { .bearer_auth(&self.config.api_key) .json(&request), provider_name, + Some(self.config.subtype.req_id_header()), ) .await?; let req_id = utils::extract_header(&raw_response, self.config.subtype.req_id_header()); @@ -199,6 +200,7 @@ impl LlmProvider for OpenAIProvider { .bearer_auth(&self.config.api_key) .json(&request), provider_name, + Some(self.config.subtype.req_id_header()), ) .await?; let req_id = utils::extract_header(&response, self.config.subtype.req_id_header()); diff --git a/server-new/src/llm/providers/utils.rs b/server-new/src/llm/providers/utils.rs index 159f9d8..343ce12 100644 --- a/server-new/src/llm/providers/utils.rs +++ b/server-new/src/llm/providers/utils.rs @@ -60,17 +60,19 @@ pub fn get_json_events( pub async fn llm_api_request( request: reqwest::RequestBuilder, provider_name: &str, + req_id_header: Option<&str>, ) -> Result { - let response = request - .send() - .await - .map_err(|e| LlmRequestError::Provider(format!("{provider_name} request failed: {e}")))?; + let response = request.send().await.map_err(|e| { + LlmRequestError::Provider(format!("{provider_name} request failed: {e}"), None) + })?; if !response.status().is_success() { let status = response.status(); + let request_id = req_id_header.and_then(|header| extract_header(&response, header)); let error_text = response.text().await.unwrap_or_default(); - return Err(LlmRequestError::Provider(format!( - "{provider_name} API error status {status}: {error_text}", - ))); + return Err(LlmRequestError::Provider( + format!("{provider_name} API error {status}: {error_text}",), + request_id, + )); } Ok(response) diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index 58c949f..505c1ce 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -124,6 +124,7 @@ impl<'r> ChatService<'r> { .id(log_id) .status(ChatRsLogStatus::Error) .error(&err.to_string()) + .maybe_request_id(err.req_id()) .build() .await?; return Err(ChatError::Request(err)); @@ -292,6 +293,7 @@ impl<'r> ChatService<'r> { .id(log_id) .status(ChatRsLogStatus::Error) .error(&err.to_string()) + .maybe_request_id(err.req_id()) .build() .await?; return Err(ChatError::Request(err)); diff --git a/server-new/src/services/chat/titles.rs b/server-new/src/services/chat/titles.rs index cd8e5ff..b1c72a2 100644 --- a/server-new/src/services/chat/titles.rs +++ b/server-new/src/services/chat/titles.rs @@ -100,6 +100,7 @@ async fn generate( .id(log_id) .status(ChatRsLogStatus::Error) .error(&err.to_string()) + .maybe_request_id(err.req_id()) .build() .await?; } From 26a1c594c6b4b4f4976700308e1d3abbf033c7a9 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 06:58:35 -0400 Subject: [PATCH 095/111] ci: update --- .github/workflows/server.yaml | 33 ++++++++++++--------------------- 1 file changed, 12 insertions(+), 21 deletions(-) diff --git a/.github/workflows/server.yaml b/.github/workflows/server.yaml index e4acbf3..53c53a6 100644 --- a/.github/workflows/server.yaml +++ b/.github/workflows/server.yaml @@ -5,6 +5,7 @@ on: tags-ignore: ["v*"] paths: - "server/**" + - "server-new/**" - "web/**" - Dockerfile pull_request: @@ -12,7 +13,7 @@ on: env: REGISTRY: ghcr.io IMAGE_NAME: ${{ github.repository }} - RUST_VERSION: "1.90" + RUST_VERSION: "1.96" NODE_VERSION: 22 jobs: @@ -23,14 +24,14 @@ jobs: run: working-directory: ./web steps: - - uses: actions/checkout@v4 - - uses: pnpm/action-setup@v4 + - uses: actions/checkout@v6 + - uses: pnpm/action-setup@v6 name: Install pnpm with: package_json_file: web/package.json run_install: false - name: Install Node.js - uses: actions/setup-node@v4 + uses: actions/setup-node@v6 with: node-version: ${{ env.NODE_VERSION }} cache: "pnpm" @@ -46,35 +47,25 @@ jobs: run: pnpm build build-server: - strategy: - matrix: - os: [ubuntu-latest, macos-latest] - name: Build server on ${{ matrix.os }} - runs-on: ${{ matrix.os }} + name: Build server + runs-on: ubuntu-latest defaults: run: - working-directory: ./server + working-directory: ./server-new steps: - - uses: actions/checkout@v4 + - uses: actions/checkout@v6 - name: Install Rust toolchain run: rustup toolchain install ${{ env.RUST_VERSION }} --profile minimal --no-self-update && rustup default ${{ env.RUST_VERSION }} - - name: Setup build dependencies on macOS - if: startsWith(runner.os, 'macOS') - run: brew link --force libpq - name: Setup rust-cache uses: Swatinem/rust-cache@v2 with: - workspaces: ./server - # - name: Install nextest - # uses: taiki-e/install-action@v2 - # with: - # tool: nextest@0.9 + workspaces: ./server-new - name: Run cargo check - run: cargo check --profile ci + run: cargo check - name: Run cargo build - run: PQ_LIB_DIR="$(brew --prefix libpq)/lib" cargo build --profile ci + run: cargo build # - name: Test crate # run: cargo nextest run --all-features --profile ci From 8871761e341070d9ef899d37acf8b240906a0ad6 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 07:12:03 -0400 Subject: [PATCH 096/111] remove unnecessary derives --- server-new/src/extractors/user.rs | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/server-new/src/extractors/user.rs b/server-new/src/extractors/user.rs index 977a09d..a97fbd1 100644 --- a/server-new/src/extractors/user.rs +++ b/server-new/src/extractors/user.rs @@ -4,7 +4,6 @@ use axum::{ extract::{FromRequestParts, OptionalFromRequestParts}, http::header, }; -use serde::{Deserialize, Serialize}; use uuid::Uuid; use crate::{db::DbService, error::AppError, state::AppState}; @@ -17,7 +16,7 @@ as an extractor in route handlers: - If used as `Option`, will be `Some` if there is an active user and `None` otherwise. */ -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug)] pub struct CurrentUser { pub user_id: Uuid, } @@ -100,8 +99,6 @@ impl OperationInput for CurrentUser { operation: &mut aide::openapi::Operation, ) { let security_reqs = [(String::from(crate::api::API_KEY_SCHEME), vec![])]; - operation - .security - .push(FromIterator::from_iter(security_reqs)) + operation.security.push(security_reqs.into()); } } From a0dfce967d24672c8733a982d86b6d854068628f Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 09:22:26 -0400 Subject: [PATCH 097/111] Update chat.rs --- server-new/src/api/chat.rs | 123 +++++++++++++++++-------------------- 1 file changed, 55 insertions(+), 68 deletions(-) diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index 600939a..f864286 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -18,24 +18,12 @@ use crate::{ api_routes! { state: AppState, tag: ApiTag::Chat.into(), - POST "/prompt" => prompt, "Prompt", { - description: "Send a single prompt to a provider and get the response" - }; - GET "/sessions" => get_active_streams, "Get active chat streams", { - description: "Get the session IDs that have ongoing response streams" - }; - GET "/sessions/{session_id}" => connect_chat_stream, "Access active chat stream", { - description: "Get a URL and token to access the response stream for this session" - }; - POST "/sessions/{session_id}" => chat_stream, "Stream chat", { - description: "Send a message in a chat session and stream the response" - }; - POST "/sessions/{session_id}/cancel" => cancel_chat_stream, "Cancel active stream", { - description: "Cancel an ongoing chat stream" - }; - POST "/session/{session_id}/regenerate" => regenerate_response, "Regenerate chat response", { - description: "Regenerate the latest assistant response in a chat session" - }; + POST "/prompt" => prompt, "Prompt"; + GET "/sessions" => get_active_streams, "Get active chat streams"; + GET "/sessions/{session_id}" => connect_chat_stream, "Access active chat stream"; + POST "/sessions/{session_id}" => chat_stream, "Stream chat session response"; + POST "/sessions/{session_id}/cancel" => cancel_chat_stream, "Cancel active chat stream"; + POST "/sessions/{session_id}/regenerate" => regenerate_response, "Regenerate chat response"; } async fn get_active_streams( @@ -50,45 +38,27 @@ async fn get_active_streams( Ok(Json(ActiveStreamsResponse { sessions })) } -#[derive(Debug, JsonSchema, serde::Serialize)] -struct ActiveStreamsResponse { - /// The chat session IDs that have ongoing response streams - sessions: Vec, -} - -#[derive(Debug, Deserialize, JsonSchema)] -struct PromptInput { - /// The prompt to send to the LLM provider - message: String, - /// The ID of the provider to chat with - provider_id: i32, - /// Configuration for the provider - options: LlmChatOptions, -} - async fn prompt( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, State(state): State, - Json(input): Json, + Json(PromptInput { + message, + provider_id, + options, + }): Json, ) -> AppResult> { let llm_provider = state .provider_service() - .build_llm_provider(&mut db, &user_id, input.provider_id) + .build_llm_provider(&mut db, &user_id, provider_id) .await?; + let prompt = LlmUserMessage { + text: message, + ..Default::default() + }; let stream_access = state .chat_service() - .prompt( - &mut db, - user_id, - input.provider_id, - llm_provider, - LlmUserMessage { - text: input.message, - ..Default::default() - }, - input.options, - ) + .prompt(&mut db, user_id, provider_id, llm_provider, prompt, options) .await?; Ok(Json(StreamAccess { @@ -97,16 +67,6 @@ async fn prompt( })) } -#[derive(Debug, Deserialize, JsonSchema)] -struct ChatInput { - /// The new chat message from the user - message: Option, - /// The ID of the provider to chat with - provider_id: i32, - /// Configuration for the provider - options: LlmChatOptions, -} - async fn chat_stream( CurrentUser { user_id }: CurrentUser, Path(session_id): Path, @@ -118,6 +78,9 @@ async fn chat_stream( .provider_service() .build_llm_provider(&mut db, &user_id, input.provider_id) .await?; + let user_message = input + .message + .map(|text| LlmUserMessage { text, files: None }); let stream_access = state .chat_service() .stream_user_chat( @@ -126,9 +89,7 @@ async fn chat_stream( session_id, input.provider_id, llm_provider, - input - .message - .map(|text| LlmUserMessage { text, files: None }), + user_message, input.options, ) .await?; @@ -139,14 +100,6 @@ async fn chat_stream( })) } -#[derive(Debug, Deserialize, JsonSchema)] -struct RegenerateInput { - /// The ID of the provider to chat with - provider_id: i32, - /// Configuration for the provider - options: LlmChatOptions, -} - async fn regenerate_response( CurrentUser { user_id }: CurrentUser, Path(session_id): Path, @@ -205,6 +158,40 @@ pub async fn cancel_chat_stream( Ok(()) } +#[derive(Debug, Deserialize, JsonSchema)] +struct PromptInput { + /// The prompt to send to the LLM provider + message: String, + /// The ID of the provider to chat with + provider_id: i32, + /// Configuration for the provider + options: LlmChatOptions, +} + +#[derive(Debug, Deserialize, JsonSchema)] +struct ChatInput { + /// The new chat message from the user + message: Option, + /// The ID of the provider to chat with + provider_id: i32, + /// Configuration for the provider + options: LlmChatOptions, +} + +#[derive(Debug, Deserialize, JsonSchema)] +struct RegenerateInput { + /// The ID of the provider to chat with + provider_id: i32, + /// Configuration for the provider + options: LlmChatOptions, +} + +#[derive(Debug, JsonSchema, serde::Serialize)] +struct ActiveStreamsResponse { + /// The chat session IDs that have ongoing response streams + sessions: Vec, +} + /// Access to an active streaming response #[derive(Serialize, JsonSchema)] struct StreamAccess { From 05a637ab4a5c55bd65915644614ac93f7098ce6a Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 10:27:01 -0400 Subject: [PATCH 098/111] use i32 for token count for consistency --- server-new/src/api/provider.rs | 6 +++--- server-new/src/db/repositories/log.rs | 8 ++------ server-new/src/llm/providers/anthropic/response.rs | 4 ++-- server-new/src/llm/providers/lorem.rs | 2 +- server-new/src/llm/providers/ollama/response.rs | 8 ++++---- server-new/src/llm/providers/openai/response.rs | 4 ++-- server-new/src/llm/types.rs | 4 ++-- 7 files changed, 16 insertions(+), 20 deletions(-) diff --git a/server-new/src/api/provider.rs b/server-new/src/api/provider.rs index 3ade43a..00ae962 100644 --- a/server-new/src/api/provider.rs +++ b/server-new/src/api/provider.rs @@ -20,10 +20,10 @@ api_routes! { state: AppState, tag: ApiTag::Provider.into(), GET "/" => list_providers, "List providers"; - GET "/{id}/models" => list_models, "List models"; + GET "/{provider_id}/models" => list_models, "List models"; POST "/" => create_provider, "Create provider"; - PATCH "/{id}" => update_provider, "Update provider"; - DELETE "/{id}" => delete_provider, "Delete provider"; + PATCH "/{provider_id}" => update_provider, "Update provider"; + DELETE "/{provider_id}" => delete_provider, "Delete provider"; } async fn list_providers( diff --git a/server-new/src/db/repositories/log.rs b/server-new/src/db/repositories/log.rs index 36ae70f..966c248 100644 --- a/server-new/src/db/repositories/log.rs +++ b/server-new/src/db/repositories/log.rs @@ -64,12 +64,8 @@ impl<'a> LogRepository<'a> { let update_log = UpdateChatRsLog { message_id, request_id, - input_tokens: usage - .and_then(|u| u.input_tokens) - .and_then(|t| t.try_into().ok()), - output_tokens: usage - .and_then(|u| u.output_tokens) - .and_then(|t| t.try_into().ok()), + input_tokens: usage.and_then(|u| u.input_tokens), + output_tokens: usage.and_then(|u| u.output_tokens), cost: usage.and_then(|u| u.cost.and_then(BigDecimal::from_f32)), error, status: status.as_ref(), diff --git a/server-new/src/llm/providers/anthropic/response.rs b/server-new/src/llm/providers/anthropic/response.rs index 804acbd..1d9bd66 100644 --- a/server-new/src/llm/providers/anthropic/response.rs +++ b/server-new/src/llm/providers/anthropic/response.rs @@ -83,8 +83,8 @@ pub enum AnthropicResponseContentBlock { /// Anthropic API response usage #[derive(Debug, Deserialize)] pub struct AnthropicUsage { - input_tokens: Option, - output_tokens: Option, + input_tokens: Option, + output_tokens: Option, } impl From for LlmUsage { diff --git a/server-new/src/llm/providers/lorem.rs b/server-new/src/llm/providers/lorem.rs index a1cc1a9..633ac7e 100644 --- a/server-new/src/llm/providers/lorem.rs +++ b/server-new/src/llm/providers/lorem.rs @@ -61,7 +61,7 @@ impl LlmProvider for LoremProvider { let response = LlmResponse { text: "Lorem ipsum".into(), usage: LlmUsage { - input_tokens: Some((prompt.text.len() / 4) as u32), + input_tokens: Some((prompt.text.len() / 4) as i32), output_tokens: Some(4), ..Default::default() }, diff --git a/server-new/src/llm/providers/ollama/response.rs b/server-new/src/llm/providers/ollama/response.rs index dd1562f..7f73996 100644 --- a/server-new/src/llm/providers/ollama/response.rs +++ b/server-new/src/llm/providers/ollama/response.rs @@ -46,11 +46,11 @@ pub struct OllamaStreamEvent { // #[serde(default)] // pub load_duration: Option, #[serde(default)] - pub prompt_eval_count: Option, + pub prompt_eval_count: Option, // #[serde(default)] // pub prompt_eval_duration: Option, #[serde(default)] - pub eval_count: Option, + pub eval_count: Option, // #[serde(default)] // pub eval_duration: Option, } @@ -73,9 +73,9 @@ pub struct OllamaCompletionResponse { // #[serde(default)] // pub eval_duration: Option, #[serde(default)] - pub prompt_eval_count: Option, + pub prompt_eval_count: Option, #[serde(default)] - pub eval_count: Option, + pub eval_count: Option, } /// Ollama message in response diff --git a/server-new/src/llm/providers/openai/response.rs b/server-new/src/llm/providers/openai/response.rs index 9c95515..1c82819 100644 --- a/server-new/src/llm/providers/openai/response.rs +++ b/server-new/src/llm/providers/openai/response.rs @@ -160,8 +160,8 @@ pub struct OpenAIResponseDelta { /// OpenAI API response usage #[derive(Debug, Deserialize)] pub struct OpenAIUsage { - prompt_tokens: Option, - completion_tokens: Option, + prompt_tokens: Option, + completion_tokens: Option, /// OpenRouter cost cost: Option, /// LLM Gateway cost diff --git a/server-new/src/llm/types.rs b/server-new/src/llm/types.rs index 40fa6dd..3d99635 100644 --- a/server-new/src/llm/types.rs +++ b/server-new/src/llm/types.rs @@ -61,8 +61,8 @@ pub struct LlmAssistantMessage { /// Usage stats from the LLM provider #[derive(Debug, Default, Clone, Copy, Serialize, Deserialize, JsonSchema)] pub struct LlmUsage { - pub input_tokens: Option, - pub output_tokens: Option, + pub input_tokens: Option, + pub output_tokens: Option, #[serde(skip_serializing_if = "Option::is_none")] pub cost: Option, } From 85933fbc36d4644ccce967924b3989c2b1f128e5 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Wed, 15 Jul 2026 23:43:47 -0400 Subject: [PATCH 099/111] llm request log improvements --- .../up.sql | 7 ++- server-new/src/api/session.rs | 37 +++++++++--- server-new/src/db/models/log.rs | 47 ++++++++++++++-- server-new/src/db/repositories/chat.rs | 56 +++++++++++-------- server-new/src/db/repositories/log.rs | 12 ++-- server-new/src/db/schema.rs | 1 + server-new/src/services/chat/mod.rs | 29 +++++----- server-new/src/services/chat/titles.rs | 14 +++-- 8 files changed, 138 insertions(+), 65 deletions(-) diff --git a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql index ca168f9..337ccf9 100644 --- a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql +++ b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql @@ -10,8 +10,9 @@ CREATE TABLE llm_logs ( input_tokens integer, output_tokens integer, cost numeric(12, 6), - status text NOT NULL, -- completed, failed, cancelled + status text NOT NULL, -- started, completed, failed, cancelled error text, + meta jsonb, -- max_tokens, temperature, etc. started_at timestamptz NOT NULL DEFAULT now(), completed_at timestamptz ); @@ -20,6 +21,8 @@ CREATE INDEX llm_logs_user_id_started_at_idx ON llm_logs (user_id, started_at DE CREATE INDEX llm_logs_session_id_idx ON llm_logs (session_id); -CREATE INDEX llm_logs_message_id_idx ON llm_logs (message_id); +CREATE UNIQUE INDEX llm_logs_message_id_unique_idx ON llm_logs (message_id) +WHERE + message_id IS NOT NULL; CREATE INDEX llm_logs_provider_id_started_at_idx ON llm_logs (provider_id, started_at DESC); diff --git a/server-new/src/api/session.rs b/server-new/src/api/session.rs index 3eb368e..0b1eff2 100644 --- a/server-new/src/api/session.rs +++ b/server-new/src/api/session.rs @@ -12,7 +12,10 @@ use uuid::Uuid; use crate::{ api::ApiTag, db::{ - models::{ChatRsMessage, ChatRsSession, NewChatRsSession, UpdateChatRsSession}, + models::{ + ChatRsLogLlmRequest, ChatRsMessage, ChatRsSession, NewChatRsSession, + UpdateChatRsSession, + }, queries::FullTextSearchResult, }, error::{AppError, AppResult}, @@ -62,15 +65,23 @@ async fn get_session( Database(mut db): Database, Path(session_id): Path, ) -> AppResult> { - let (session, messages) = db + let session = db .chats() - .find_session_with_messages(&user_id, &session_id) - .await?; + .find_session(&user_id, &session_id) + .await? + .ok_or_else(|| AppError::not_found("chat session not found"))?; + let messages = db + .chats() + .list_messages_with_logs(&session_id) + .await? + .into_iter() + .map(|(message, llm_request)| SessionMessage { + message, + llm_request, + }) + .collect(); - match session { - Some(session) => Ok(Json(GetSessionResponse { session, messages })), - None => Err(AppError::not_found("chat session not found")), - } + Ok(Json(GetSessionResponse { session, messages })) } async fn search_sessions( @@ -145,7 +156,15 @@ struct MessageIdResponse { #[derive(Serialize, JsonSchema)] struct GetSessionResponse { session: ChatRsSession, - messages: Vec, + messages: Vec, +} + +#[derive(Serialize, JsonSchema)] +struct SessionMessage { + /// The message + message: ChatRsMessage, + /// Request metadata for assistant responses + llm_request: Option, } #[derive(Deserialize, JsonSchema)] diff --git a/server-new/src/db/models/log.rs b/server-new/src/db/models/log.rs index 914f2b3..caec46b 100644 --- a/server-new/src/db/models/log.rs +++ b/server-new/src/db/models/log.rs @@ -1,14 +1,20 @@ use bigdecimal::BigDecimal; -use diesel::prelude::*; +use diesel::{deserialize::FromSqlRow, expression::AsExpression, prelude::*}; +use diesel_jsonb_derive::AsJsonb; use schemars::JsonSchema; -use serde::Serialize; +use serde::{Deserialize, Serialize}; +use serde_with::skip_serializing_none; use strum::{AsRefStr, EnumString}; use uuid::Uuid; -use crate::db::{UtcDateTime, models::ChatRsUser}; +use crate::db::{ + UtcDateTime, + models::{ChatRsMessage, ChatRsUser}, +}; -#[derive(Identifiable, Associations, Queryable, Selectable, Serialize, JsonSchema)] +#[derive(Identifiable, Associations, Queryable, Selectable)] #[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] +#[diesel(belongs_to(ChatRsMessage, foreign_key = message_id))] #[diesel(table_name = super::schema::llm_logs)] pub struct ChatRsLog { pub id: i32, @@ -24,11 +30,13 @@ pub struct ChatRsLog { pub cost: Option, pub status: String, pub error: Option, + pub meta: Option, pub started_at: UtcDateTime, pub completed_at: Option, } -#[derive(Debug, Clone, Copy, EnumString, AsRefStr)] +#[derive(Debug, Clone, Copy, EnumString, AsRefStr, Serialize, JsonSchema)] +#[serde(rename_all = "lowercase")] #[strum(serialize_all = "lowercase")] pub enum ChatRsLogKind { Chat, @@ -37,7 +45,8 @@ pub enum ChatRsLogKind { Image, } -#[derive(Debug, Clone, Copy, EnumString, AsRefStr)] +#[derive(Debug, Clone, Copy, EnumString, AsRefStr, Serialize, JsonSchema)] +#[serde(rename_all = "lowercase")] #[strum(serialize_all = "lowercase")] pub enum ChatRsLogStatus { Started, @@ -46,6 +55,31 @@ pub enum ChatRsLogStatus { Error, } +#[skip_serializing_none] +#[derive(Debug, Default, Serialize, Deserialize, FromSqlRow, AsExpression, AsJsonb)] +#[diesel(sql_type = diesel::sql_types::Jsonb)] +pub struct ChatRsLogMeta { + pub temperature: Option, + pub max_tokens: Option, +} + +#[skip_serializing_none] +#[derive(Identifiable, Queryable, Selectable, Serialize, JsonSchema)] +#[diesel(table_name = super::schema::llm_logs)] +pub struct ChatRsLogLlmRequest { + #[serde(skip)] + pub id: i32, + pub provider_id: Option, + pub model: String, + pub request_id: Option, + pub input_tokens: Option, + pub output_tokens: Option, + pub cost: Option, + #[schemars(with = "ChatRsLogStatus")] + pub status: String, + pub error: Option, +} + #[derive(Insertable)] #[diesel(table_name = super::schema::llm_logs)] pub struct NewChatRsLog<'a> { @@ -55,6 +89,7 @@ pub struct NewChatRsLog<'a> { pub session_id: Option<&'a Uuid>, pub model: &'a str, pub status: &'a str, + pub meta: Option<&'a ChatRsLogMeta>, } #[derive(Default, AsChangeset)] diff --git a/server-new/src/db/repositories/chat.rs b/server-new/src/db/repositories/chat.rs index f61d73f..f9d835b 100644 --- a/server-new/src/db/repositories/chat.rs +++ b/server-new/src/db/repositories/chat.rs @@ -5,10 +5,11 @@ use uuid::Uuid; use crate::db::{ DbConnection, models::{ - ChatRsMessage, ChatRsSession, NewChatRsMessage, NewChatRsSession, UpdateChatRsSession, + ChatRsLogKind, ChatRsLogLlmRequest, ChatRsMessage, ChatRsSession, NewChatRsMessage, + NewChatRsSession, UpdateChatRsSession, }, queries::{FullTextSearchResult, full_text_query}, - schema::{chat_messages, chat_sessions}, + schema::{chat_messages, chat_sessions, llm_logs}, }; pub struct ChatRepository<'a> { @@ -116,30 +117,37 @@ impl<'a> ChatRepository<'a> { Ok(session) } - pub async fn find_session_with_messages( + pub async fn list_messages(&mut self, session_id: &Uuid) -> QueryResult> { + let messages = chat_messages::table + .filter(chat_messages::session_id.eq(session_id)) + .select(ChatRsMessage::as_select()) + .order_by(chat_messages::created_at.asc()) + .load(self.db) + .await?; + + Ok(messages) + } + + pub async fn list_messages_with_logs( &mut self, - user_id: &Uuid, session_id: &Uuid, - ) -> Result<(Option, Vec), diesel::result::Error> { - let (session, messages) = futures::future::join( - chat_sessions::table - .filter(chat_sessions::user_id.eq(user_id)) - .filter(chat_sessions::id.eq(session_id)) - .select(ChatRsSession::as_select()) - .first(&mut &**self.db), - chat_messages::table - .inner_join( - chat_sessions::table.on(chat_sessions::id.eq(chat_messages::session_id)), - ) - .filter(chat_sessions::user_id.eq(user_id)) - .filter(chat_messages::session_id.eq(session_id)) - .select(ChatRsMessage::as_select()) - .order_by(chat_messages::created_at.asc()) - .load(&mut &**self.db), - ) - .await; - - Ok((session.optional()?, messages?)) + ) -> QueryResult)>> { + let messages = chat_messages::table + .left_join( + llm_logs::table.on(llm_logs::message_id + .eq(chat_messages::id.nullable()) + .and(llm_logs::kind.eq(ChatRsLogKind::Chat.as_ref()))), + ) + .filter(chat_messages::session_id.eq(session_id)) + .select(( + ChatRsMessage::as_select(), + Option::::as_select(), + )) + .order_by(chat_messages::created_at.asc()) + .load(self.db) + .await?; + + Ok(messages) } pub async fn search_sessions( diff --git a/server-new/src/db/repositories/log.rs b/server-new/src/db/repositories/log.rs index 966c248..414ed15 100644 --- a/server-new/src/db/repositories/log.rs +++ b/server-new/src/db/repositories/log.rs @@ -7,10 +7,10 @@ use uuid::Uuid; use crate::{ db::{ DbConnection, UtcDateTime, - models::{ChatRsLogKind, ChatRsLogStatus, NewChatRsLog, UpdateChatRsLog}, + models::{ChatRsLogKind, ChatRsLogMeta, ChatRsLogStatus, NewChatRsLog, UpdateChatRsLog}, schema::llm_logs, }, - llm::types::LlmUsage, + llm::types::{LlmChatOptions, LlmUsage}, }; pub struct LogRepository<'a> { @@ -29,7 +29,7 @@ impl<'a> LogRepository<'a> { &mut self, user_id: &Uuid, provider_id: i32, - model: &str, + llm_options: &LlmChatOptions, kind: ChatRsLogKind, session_id: Option<&Uuid>, ) -> QueryResult { @@ -38,8 +38,12 @@ impl<'a> LogRepository<'a> { user_id, provider_id, session_id, - model, + model: &llm_options.model, status: ChatRsLogStatus::Started.as_ref(), + meta: Some(&ChatRsLogMeta { + temperature: llm_options.temperature, + max_tokens: llm_options.max_tokens, + }), }; diesel::insert_into(llm_logs::table) diff --git a/server-new/src/db/schema.rs b/server-new/src/db/schema.rs index 4858351..5622f08 100644 --- a/server-new/src/db/schema.rs +++ b/server-new/src/db/schema.rs @@ -99,6 +99,7 @@ diesel::table! { cost -> Nullable, status -> Text, error -> Nullable, + meta -> Nullable, started_at -> Timestamptz, completed_at -> Nullable, } diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index 505c1ce..c95cc55 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -108,7 +108,7 @@ impl<'r> ChatService<'r> { .kind(ChatRsLogKind::Prompt) .user_id(&user_id) .provider_id(provider_id) - .model(&options.model) + .llm_options(&options) .build() .await?; @@ -173,11 +173,12 @@ impl<'r> ChatService<'r> { user_message: Option, chat_options: LlmChatOptions, ) -> Result { - let (chat_session, mut messages) = db - .chats() - .find_session_with_messages(&user_id, &session_id) - .await?; - let chat_session = chat_session.ok_or(ChatError::SessionNotFound)?; + let mut chats = db.chats(); + let chat_session = chats + .find_session(&user_id, &session_id) + .await? + .ok_or(ChatError::SessionNotFound)?; + let mut messages = chats.list_messages(&session_id).await?; if let Some(user_message) = user_message { if messages.is_empty() && chat_session.title == DEFAULT_SESSION_TITLE { @@ -191,8 +192,7 @@ impl<'r> ChatService<'r> { self.db_pool, ); } - let new_message = db - .chats() + let new_message = chats .save_message(NewChatRsMessage { content: &user_message.text, session_id: &session_id, @@ -228,11 +228,12 @@ impl<'r> ChatService<'r> { provider: Arc, chat_options: LlmChatOptions, ) -> Result { - let (chat_session, mut messages) = db - .chats() - .find_session_with_messages(&user_id, &session_id) - .await?; - chat_session.ok_or(ChatError::SessionNotFound)?; + let mut chats = db.chats(); + let chat_session = chats + .find_session(&user_id, &session_id) + .await? + .ok_or(ChatError::SessionNotFound)?; + let mut messages = chats.list_messages(&chat_session.id).await?; let last_message = messages.pop(); if last_message.as_ref().is_none_or(|m| !m.role.is_assistant()) { @@ -275,7 +276,7 @@ impl<'r> ChatService<'r> { .user_id(¶ms.user_id) .session_id(¶ms.session_id) .provider_id(params.provider_id) - .model(¶ms.chat_options.model) + .llm_options(¶ms.chat_options) .build() .await?; diff --git a/server-new/src/services/chat/titles.rs b/server-new/src/services/chat/titles.rs index b1c72a2..92816c8 100644 --- a/server-new/src/services/chat/titles.rs +++ b/server-new/src/services/chat/titles.rs @@ -51,6 +51,12 @@ async fn generate( model: String, db_pool: DbPool, ) -> Result<(), ChatError> { + let prompt_options = LlmChatOptions { + model, + temperature: Some(TITLE_PROMPT_TEMPERATURE), + max_tokens: Some(TITLE_PROMPT_MAX_TOKENS), + }; + let mut db = DbService::from_pool(&db_pool).await?; let log_id = db .logs() @@ -59,17 +65,13 @@ async fn generate( .session_id(&session_id) .provider_id(provider_id) .kind(ChatRsLogKind::Title) - .model(&model) + .llm_options(&prompt_options) .build() .await?; let prompt = LlmPrompt { text: &format!("{TITLE_PROMPT}\n\n\"{user_message}\""), - options: &LlmChatOptions { - model, - temperature: Some(TITLE_PROMPT_TEMPERATURE), - max_tokens: Some(TITLE_PROMPT_MAX_TOKENS), - }, + options: &prompt_options, }; match provider.prompt(prompt).await { Ok(LlmResponse { text, usage, meta }) => { From 81cb6574940b1ada699c3a661e2b92be42aa5707 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Thu, 16 Jul 2026 13:08:55 -0400 Subject: [PATCH 100/111] add ttft --- server-new/Cargo.lock | 1 - server-new/Cargo.toml | 1 - .../up.sql | 5 +- server-new/src/db/mod.rs | 2 +- server-new/src/db/models/log.rs | 66 +++++----- server-new/src/db/repositories.rs | 2 +- server-new/src/db/repositories/log.rs | 93 ++++++++----- server-new/src/db/schema.rs | 5 +- .../src/llm/providers/anthropic/response.rs | 3 + server-new/src/services/chat/mod.rs | 122 +++++++++--------- server-new/src/services/chat/titles.rs | 55 ++++---- server-new/src/services/stream/mod.rs | 8 +- server-new/src/services/stream/writer.rs | 31 +++-- 13 files changed, 223 insertions(+), 171 deletions(-) diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index ec59900..2591421 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -2439,7 +2439,6 @@ dependencies = [ "axum-helmet", "axum-plugin", "bigdecimal", - "bon", "chrono", "diesel", "diesel-async", diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index b0e46c4..91ab789 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -27,7 +27,6 @@ axum-plugin = { features = ["figment"] } bigdecimal = { version = "0.4.10", features = ["serde-json"] } -bon = "3.9.3" chrono = { version = "0.4.45", default-features = false, diff --git a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql index 337ccf9..850a878 100644 --- a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql +++ b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql @@ -6,13 +6,12 @@ CREATE TABLE llm_logs ( session_id uuid REFERENCES chat_sessions (id) ON DELETE SET NULL, message_id uuid REFERENCES chat_messages (id) ON DELETE SET NULL, model text NOT NULL, - request_id text, -- provider request ID input_tokens integer, output_tokens integer, cost numeric(12, 6), + ttft_ms integer, -- time to first token in milliseconds status text NOT NULL, -- started, completed, failed, cancelled - error text, - meta jsonb, -- max_tokens, temperature, etc. + meta jsonb NOT NULL DEFAULT '{}', -- max_tokens, temperature, etc. started_at timestamptz NOT NULL DEFAULT now(), completed_at timestamptz ); diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs index a6def9b..a448e87 100644 --- a/server-new/src/db/mod.rs +++ b/server-new/src/db/mod.rs @@ -9,7 +9,7 @@ use diesel_async::{ pub mod models; pub mod queries; -mod repositories; +pub mod repositories; mod schema; /// Type of the database pool diff --git a/server-new/src/db/models/log.rs b/server-new/src/db/models/log.rs index caec46b..50c1d85 100644 --- a/server-new/src/db/models/log.rs +++ b/server-new/src/db/models/log.rs @@ -12,7 +12,7 @@ use crate::db::{ models::{ChatRsMessage, ChatRsUser}, }; -#[derive(Identifiable, Associations, Queryable, Selectable)] +#[derive(Identifiable, Associations, Queryable, Selectable, AsChangeset)] #[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] #[diesel(belongs_to(ChatRsMessage, foreign_key = message_id))] #[diesel(table_name = super::schema::llm_logs)] @@ -24,60 +24,62 @@ pub struct ChatRsLog { pub session_id: Option, pub message_id: Option, pub model: String, - pub request_id: Option, pub input_tokens: Option, pub output_tokens: Option, pub cost: Option, + pub ttft_ms: Option, pub status: String, - pub error: Option, - pub meta: Option, + pub meta: ChatRsLogMeta, pub started_at: UtcDateTime, pub completed_at: Option, } -#[derive(Debug, Clone, Copy, EnumString, AsRefStr, Serialize, JsonSchema)] +#[skip_serializing_none] +#[derive(Identifiable, Queryable, Selectable, Serialize, JsonSchema)] +#[diesel(table_name = super::schema::llm_logs)] +pub struct ChatRsLogLlmRequest { + #[serde(skip)] + pub id: i32, + pub provider_id: Option, + pub model: String, + pub input_tokens: Option, + pub output_tokens: Option, + pub cost: Option, + #[schemars(with = "ChatRsLogStatus")] + pub status: String, + pub meta: ChatRsLogMeta, +} + +#[derive(Debug, Default, Clone, Copy, EnumString, AsRefStr, Serialize, JsonSchema)] #[serde(rename_all = "lowercase")] #[strum(serialize_all = "lowercase")] pub enum ChatRsLogKind { Chat, Title, + #[default] Prompt, Image, } -#[derive(Debug, Clone, Copy, EnumString, AsRefStr, Serialize, JsonSchema)] +#[derive(Debug, Default, Clone, Copy, EnumString, AsRefStr, Serialize, JsonSchema)] #[serde(rename_all = "lowercase")] #[strum(serialize_all = "lowercase")] pub enum ChatRsLogStatus { Started, - Completed, Cancelled, Error, + #[default] + Completed, } #[skip_serializing_none] -#[derive(Debug, Default, Serialize, Deserialize, FromSqlRow, AsExpression, AsJsonb)] +#[derive(Debug, Default, Serialize, Deserialize, FromSqlRow, AsExpression, AsJsonb, JsonSchema)] #[diesel(sql_type = diesel::sql_types::Jsonb)] pub struct ChatRsLogMeta { pub temperature: Option, pub max_tokens: Option, -} - -#[skip_serializing_none] -#[derive(Identifiable, Queryable, Selectable, Serialize, JsonSchema)] -#[diesel(table_name = super::schema::llm_logs)] -pub struct ChatRsLogLlmRequest { - #[serde(skip)] - pub id: i32, - pub provider_id: Option, - pub model: String, + pub errors: Option>, pub request_id: Option, - pub input_tokens: Option, - pub output_tokens: Option, - pub cost: Option, - #[schemars(with = "ChatRsLogStatus")] - pub status: String, - pub error: Option, } #[derive(Insertable)] @@ -90,17 +92,19 @@ pub struct NewChatRsLog<'a> { pub model: &'a str, pub status: &'a str, pub meta: Option<&'a ChatRsLogMeta>, + pub started_at: UtcDateTime, } -#[derive(Default, AsChangeset)] +#[derive(Default, Identifiable, Queryable, Selectable, AsChangeset)] #[diesel(table_name = super::schema::llm_logs)] -pub struct UpdateChatRsLog<'a> { - pub message_id: Option<&'a Uuid>, - pub request_id: Option<&'a str>, +pub struct UpdateChatRsLog { + pub id: i32, + pub message_id: Option, pub input_tokens: Option, pub output_tokens: Option, + pub ttft_ms: Option, pub cost: Option, - pub status: &'a str, - pub error: Option<&'a str>, - pub completed_at: UtcDateTime, + pub status: String, + pub meta: ChatRsLogMeta, + pub completed_at: Option, } diff --git a/server-new/src/db/repositories.rs b/server-new/src/db/repositories.rs index f592d53..31aac27 100644 --- a/server-new/src/db/repositories.rs +++ b/server-new/src/db/repositories.rs @@ -10,7 +10,7 @@ mod user; pub use api_key::ApiKeyRepository; pub use chat::ChatRepository; -pub use log::LogRepository; +pub use log::{LlmLogComplete, LlmLogCreate, LogRepository}; pub use provider::ProviderRepository; pub use secret::SecretRepository; pub use session::SessionRepository; diff --git a/server-new/src/db/repositories/log.rs b/server-new/src/db/repositories/log.rs index 414ed15..103a4d9 100644 --- a/server-new/src/db/repositories/log.rs +++ b/server-new/src/db/repositories/log.rs @@ -1,5 +1,6 @@ +use std::time::Duration; + use bigdecimal::{BigDecimal, FromPrimitive}; -use bon::bon; use diesel::prelude::*; use diesel_async::RunQueryDsl; use uuid::Uuid; @@ -16,71 +17,101 @@ use crate::{ pub struct LogRepository<'a> { db: &'a mut DbConnection, } - -#[bon] impl<'a> LogRepository<'a> { pub fn new(db: &'a mut DbConnection) -> Self { LogRepository { db } } /// Create a new LLM request log entry - #[builder(finish_fn = "build")] pub async fn create( &mut self, - user_id: &Uuid, - provider_id: i32, - llm_options: &LlmChatOptions, - kind: ChatRsLogKind, - session_id: Option<&Uuid>, - ) -> QueryResult { + LlmLogCreate { + user_id, + provider_id, + llm_options, + kind, + session_id, + }: LlmLogCreate<'_>, + ) -> QueryResult { let new_log = NewChatRsLog { kind: kind.as_ref(), - user_id, + user_id: &user_id, provider_id, session_id, - model: &llm_options.model, + model: llm_options + .as_ref() + .map(|o| o.model.as_str()) + .unwrap_or_default(), status: ChatRsLogStatus::Started.as_ref(), meta: Some(&ChatRsLogMeta { - temperature: llm_options.temperature, - max_tokens: llm_options.max_tokens, + temperature: llm_options.and_then(|o| o.temperature), + max_tokens: llm_options.and_then(|o| o.max_tokens), + ..Default::default() }), + started_at: chrono::Utc::now(), }; diesel::insert_into(llm_logs::table) .values(new_log) - .returning(llm_logs::id) + .returning(UpdateChatRsLog::as_returning()) .get_result(self.db) .await } /// Complete a LLM request log entry - #[builder(finish_fn = "build")] pub async fn complete( &mut self, - id: i32, - message_id: Option<&Uuid>, - request_id: Option<&str>, - usage: Option<&LlmUsage>, - error: Option<&str>, - status: ChatRsLogStatus, - completed_at: Option, - ) -> QueryResult { - let update_log = UpdateChatRsLog { + log: UpdateChatRsLog, + LlmLogComplete { message_id, request_id, + usage, + errors, + first_token_in, + status, + completed_at, + }: LlmLogComplete<'_>, + ) -> QueryResult { + let updated_log = UpdateChatRsLog { + id: log.id, + message_id, input_tokens: usage.and_then(|u| u.input_tokens), output_tokens: usage.and_then(|u| u.output_tokens), cost: usage.and_then(|u| u.cost.and_then(BigDecimal::from_f32)), - error, - status: status.as_ref(), - completed_at: completed_at.unwrap_or_else(chrono::Utc::now), + status: status.as_ref().to_owned(), + completed_at: Some(completed_at.unwrap_or_else(chrono::Utc::now)), + ttft_ms: first_token_in.and_then(|d| d.as_millis().try_into().ok()), + meta: ChatRsLogMeta { + request_id: request_id.map(str::to_owned), + errors, + ..log.meta + }, }; - diesel::update(llm_logs::table) - .filter(llm_logs::id.eq(id)) - .set(update_log) + diesel::update(&updated_log) + .set(&updated_log) .returning(llm_logs::id) .get_result(self.db) .await } } + +#[derive(Debug, Default)] +pub struct LlmLogCreate<'a> { + pub kind: ChatRsLogKind, + pub user_id: Uuid, + pub provider_id: i32, + pub session_id: Option<&'a Uuid>, + pub llm_options: Option<&'a LlmChatOptions>, +} + +#[derive(Debug, Default)] +pub struct LlmLogComplete<'a> { + pub status: ChatRsLogStatus, + pub message_id: Option, + pub request_id: Option<&'a str>, + pub usage: Option<&'a LlmUsage>, + pub errors: Option>, + pub first_token_in: Option, + pub completed_at: Option, +} diff --git a/server-new/src/db/schema.rs b/server-new/src/db/schema.rs index 5622f08..07aff3b 100644 --- a/server-new/src/db/schema.rs +++ b/server-new/src/db/schema.rs @@ -93,13 +93,12 @@ diesel::table! { session_id -> Nullable, message_id -> Nullable, model -> Text, - request_id -> Nullable, input_tokens -> Nullable, output_tokens -> Nullable, cost -> Nullable, + ttft_ms -> Nullable, status -> Text, - error -> Nullable, - meta -> Nullable, + meta -> Jsonb, started_at -> Timestamptz, completed_at -> Nullable, } diff --git a/server-new/src/llm/providers/anthropic/response.rs b/server-new/src/llm/providers/anthropic/response.rs index 1d9bd66..42e6f3b 100644 --- a/server-new/src/llm/providers/anthropic/response.rs +++ b/server-new/src/llm/providers/anthropic/response.rs @@ -75,6 +75,7 @@ pub fn parse_anthropic_event( /// Anthropic API response content block #[derive(Debug, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] +#[allow(unused)] pub enum AnthropicResponseContentBlock { Text { text: String }, ToolUse { id: String, name: String }, @@ -121,6 +122,7 @@ pub struct AnthropicStreamResponse { /// Anthropic streaming event types #[derive(Debug, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] +#[allow(unused)] pub enum AnthropicStreamEvent { MessageStart { message: AnthropicStreamResponse, @@ -149,6 +151,7 @@ pub enum AnthropicStreamEvent { #[derive(Debug, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] +#[allow(unused)] pub enum AnthropicDelta { TextDelta { text: String }, InputJsonDelta { partial_json: String }, diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index c95cc55..0dfe10a 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -1,4 +1,4 @@ -use std::sync::Arc; +use std::{sync::Arc, time::Instant}; use tinistream_client::types::{StreamAccessResponse, StreamStatus}; use uuid::Uuid; @@ -8,8 +8,9 @@ use crate::{ DbPool, DbService, models::{ AssistantMeta, ChatRsLogKind, ChatRsLogStatus, ChatRsMessage, ChatRsMessageMeta, - ChatRsMessageRole, NewChatRsMessage, UserMeta, + ChatRsMessageRole, NewChatRsMessage, UpdateChatRsLog, UserMeta, }, + repositories::{LlmLogComplete, LlmLogCreate}, }, llm::{ interface::{LlmProvider, LlmResponseMeta}, @@ -102,15 +103,14 @@ impl<'r> ChatService<'r> { prompt: LlmUserMessage, options: LlmChatOptions, ) -> Result { - let log_id = db - .logs() - .create() - .kind(ChatRsLogKind::Prompt) - .user_id(&user_id) - .provider_id(provider_id) - .llm_options(&options) - .build() - .await?; + let create_log = LlmLogCreate { + kind: ChatRsLogKind::Prompt, + user_id, + provider_id, + llm_options: Some(&options), + ..Default::default() + }; + let log = db.logs().create(create_log).await?; let request = LlmChatRequest { messages: &[LlmMessage::User(prompt)], @@ -119,20 +119,20 @@ impl<'r> ChatService<'r> { let (response_stream, response_meta) = match provider.stream_chat(request).await { Ok(response) => response, Err(err) => { - db.logs() - .complete() - .id(log_id) - .status(ChatRsLogStatus::Error) - .error(&err.to_string()) - .maybe_request_id(err.req_id()) - .build() - .await?; + let complete_log = LlmLogComplete { + status: ChatRsLogStatus::Error, + request_id: err.req_id(), + errors: Some(vec![err.to_string()]), + ..Default::default() + }; + db.logs().complete(log, complete_log).await?; return Err(ChatError::Request(err)); } }; // Spawn thread to process LLM streaming response let stream_key = StreamingService::prompt_key(&user_id); + let start_time = Instant::now(); let (stream_access, ws_writer, ws_reader) = StreamingService::new(self.tinistream) .create_stream(&stream_key) .await?; @@ -140,18 +140,18 @@ impl<'r> ChatService<'r> { let tinistream_client = self.tinistream.to_owned(); tokio::spawn(async move { let output = - StreamingService::process_stream(response_stream, ws_writer, ws_reader).await; - if let Ok(mut db) = DbService::from_pool(&db_pool).await { - let _ = db - .logs() - .complete() - .id(log_id) - .status(output.status()) - .maybe_request_id(response_meta.request_id.as_deref()) - .maybe_usage(output.usage.as_ref()) - .maybe_error(output.errors.map(|e| e.join(", ")).as_deref()) - .build() + StreamingService::process_stream(response_stream, start_time, ws_writer, ws_reader) .await; + if let Ok(mut db) = DbService::from_pool(&db_pool).await { + let complete_log = LlmLogComplete { + status: output.status(), + request_id: response_meta.request_id.as_deref(), + usage: output.usage.as_ref(), + errors: output.errors, + first_token_in: output.first_token_in, + ..Default::default() + }; + let _ = db.logs().complete(log, complete_log).await; } let _ = StreamingService::new(&tinistream_client) .end_stream(&stream_key) @@ -269,17 +269,16 @@ impl<'r> ChatService<'r> { return Err(ChatError::AlreadyStreaming); } - let log_id = db - .logs() - .create() - .kind(ChatRsLogKind::Chat) - .user_id(¶ms.user_id) - .session_id(¶ms.session_id) - .provider_id(params.provider_id) - .llm_options(¶ms.chat_options) - .build() - .await?; + let create_log = LlmLogCreate { + kind: ChatRsLogKind::Chat, + user_id: params.user_id, + provider_id: params.provider_id, + llm_options: Some(¶ms.chat_options), + session_id: Some(¶ms.session_id), + }; + let log = db.logs().create(create_log).await?; + let start_time = Instant::now(); let (response_stream, meta) = match provider .stream_chat(LlmChatRequest { messages: &messages::build_llm_messages(messages)?, @@ -289,14 +288,13 @@ impl<'r> ChatService<'r> { { Ok(response) => response, Err(err) => { - db.logs() - .complete() - .id(log_id) - .status(ChatRsLogStatus::Error) - .error(&err.to_string()) - .maybe_request_id(err.req_id()) - .build() - .await?; + let complete_log = LlmLogComplete { + status: ChatRsLogStatus::Error, + request_id: err.req_id(), + errors: Some(vec![err.to_string()]), + ..Default::default() + }; + db.logs().complete(log, complete_log).await?; return Err(ChatError::Request(err)); } }; @@ -307,9 +305,10 @@ impl<'r> ChatService<'r> { let tinistream_client = self.tinistream.to_owned(); tokio::spawn(async move { let output = - StreamingService::process_stream(response_stream, ws_writer, ws_reader).await; + StreamingService::process_stream(response_stream, start_time, ws_writer, ws_reader) + .await; let stream_cancelled = output.cancelled; - if let Err(err) = Self::persist_response(output, params, log_id, meta, db_pool).await { + if let Err(err) = Self::persist_response(output, params, log, meta, db_pool).await { tracing::error!("Failed to save assistant response: {err}"); } @@ -328,7 +327,7 @@ impl<'r> ChatService<'r> { async fn persist_response( output: LlmStreamOutput, params: ChatStreamParams, - log_id: i32, + log: UpdateChatRsLog, meta: LlmResponseMeta, db_pool: DbPool, ) -> Result { @@ -360,17 +359,16 @@ impl<'r> ChatService<'r> { .await?; } - db.logs() - .complete() - .id(log_id) - .status(output.status()) - .completed_at(completed_at) - .message_id(&new_message.id) - .maybe_request_id(meta.request_id.as_deref()) - .maybe_usage(output.usage.as_ref()) - .maybe_error(output.errors.map(|e| e.join(", ")).as_deref()) - .build() - .await?; + let complete_log = LlmLogComplete { + status: output.status(), + message_id: Some(new_message.id), + request_id: meta.request_id.as_deref(), + usage: output.usage.as_ref(), + errors: output.errors, + first_token_in: output.first_token_in, + completed_at: Some(completed_at), + }; + db.logs().complete(log, complete_log).await?; Ok(new_message) } diff --git a/server-new/src/services/chat/titles.rs b/server-new/src/services/chat/titles.rs index 92816c8..ff2e507 100644 --- a/server-new/src/services/chat/titles.rs +++ b/server-new/src/services/chat/titles.rs @@ -6,6 +6,7 @@ use crate::{ db::{ DbPool, DbService, models::{ChatRsLogKind, ChatRsLogStatus, UpdateChatRsSession}, + repositories::{LlmLogComplete, LlmLogCreate}, }, llm::{ interface::{LlmProvider, LlmResponse}, @@ -58,16 +59,15 @@ async fn generate( }; let mut db = DbService::from_pool(&db_pool).await?; - let log_id = db - .logs() - .create() - .user_id(&user_id) - .session_id(&session_id) - .provider_id(provider_id) - .kind(ChatRsLogKind::Title) - .llm_options(&prompt_options) - .build() - .await?; + let create_log = LlmLogCreate { + kind: ChatRsLogKind::Title, + user_id, + provider_id, + session_id: Some(&session_id), + llm_options: Some(&prompt_options), + }; + let log = db.logs().create(create_log).await?; + drop(db); let prompt = LlmPrompt { text: &format!("{TITLE_PROMPT}\n\n\"{user_message}\""), @@ -76,6 +76,8 @@ async fn generate( match provider.prompt(prompt).await { Ok(LlmResponse { text, usage, meta }) => { let completed_at = chrono::Utc::now(); + + let mut db = DbService::from_pool(&db_pool).await?; db.chats() .update_session( &user_id, @@ -86,25 +88,24 @@ async fn generate( }, ) .await?; - db.logs() - .complete() - .id(log_id) - .status(ChatRsLogStatus::Completed) - .completed_at(completed_at) - .usage(&usage) - .maybe_request_id(meta.request_id.as_deref()) - .build() - .await?; + + let complete_log = LlmLogComplete { + usage: Some(&usage), + request_id: meta.request_id.as_deref(), + completed_at: Some(completed_at), + ..Default::default() + }; + db.logs().complete(log, complete_log).await?; } Err(err) => { - db.logs() - .complete() - .id(log_id) - .status(ChatRsLogStatus::Error) - .error(&err.to_string()) - .maybe_request_id(err.req_id()) - .build() - .await?; + let mut db = DbService::from_pool(&db_pool).await?; + let complete_log = LlmLogComplete { + status: ChatRsLogStatus::Error, + request_id: err.req_id(), + errors: Some(vec![err.to_string()]), + ..Default::default() + }; + db.logs().complete(log, complete_log).await?; } } diff --git a/server-new/src/services/stream/mod.rs b/server-new/src/services/stream/mod.rs index 06da595..6f6013d 100644 --- a/server-new/src/services/stream/mod.rs +++ b/server-new/src/services/stream/mod.rs @@ -1,3 +1,5 @@ +use std::time::{Duration, Instant}; + use futures::{ StreamExt, stream::{SplitSink, SplitStream}, @@ -32,6 +34,7 @@ pub struct LlmStreamOutput { // pub images: Option>, pub usage: Option, pub errors: Option>, + pub first_token_in: Option, pub cancelled: bool, } impl LlmStreamOutput { @@ -39,7 +42,7 @@ impl LlmStreamOutput { pub fn status(&self) -> ChatRsLogStatus { if self.cancelled { ChatRsLogStatus::Cancelled - } else if self.errors.is_some() { + } else if self.errors.as_ref().is_some_and(|e| !e.is_empty()) { ChatRsLogStatus::Error } else { ChatRsLogStatus::Completed @@ -114,11 +117,12 @@ impl<'r> StreamingService<'r> { /// and return the accumulated response. pub async fn process_stream( stream: LlmStream, + start_time: Instant, writer: WsWriter, reader: WsReader, ) -> LlmStreamOutput { writer::LlmStreamWriter::new() - .process(stream, writer, reader) + .process(stream, start_time, writer, reader) .await } diff --git a/server-new/src/services/stream/writer.rs b/server-new/src/services/stream/writer.rs index 704f39a..4f49280 100644 --- a/server-new/src/services/stream/writer.rs +++ b/server-new/src/services/stream/writer.rs @@ -34,6 +34,8 @@ pub struct LlmStreamWriter { errors: Option>, /// Accumulated usage information from the LLM provider. usage: Option, + /// Duration from request start to first token + first_token_in: Option, } /// Internal state @@ -64,6 +66,7 @@ impl LlmStreamWriter { // images: None, errors: None, usage: None, + first_token_in: None, } } @@ -73,6 +76,7 @@ impl LlmStreamWriter { pub async fn process( &mut self, stream: LlmStream, + start_time: Instant, mut writer: WsWriter, mut reader: WsReader, ) -> LlmStreamOutput { @@ -90,7 +94,7 @@ impl LlmStreamWriter { }); tokio::select! { - _ = self.process_stream(stream, &mut writer) => {} + _ = self.process_stream(stream, start_time, &mut writer) => {} _ = cancel_token.cancelled() => { self.errors.get_or_insert_default().push(LlmStreamChunkError::StreamCancelled); cancelled = true; @@ -110,16 +114,27 @@ impl LlmStreamWriter { .map(|e| e.to_string()) .collect::>() }), + first_token_in: self.first_token_in.take(), cancelled, } } - async fn process_stream(&mut self, mut stream: LlmStream, writer: &mut WsWriter) { - let mut last_flush_time = Instant::now(); + async fn process_stream( + &mut self, + mut stream: LlmStream, + start_time: Instant, + writer: &mut WsWriter, + ) { + let mut last_flushed_at = Instant::now(); loop { match stream.next().await { Some(Ok(chunk)) => match chunk { - LlmStreamChunk::Text(text) => self.process_text(&text), + LlmStreamChunk::Text(text) => { + if self.first_token_in.is_none() { + self.first_token_in = Some(start_time.elapsed()); + } + self.process_text(&text); + } // LlmStreamChunk::ToolCalls(tool_calls) => self.process_tool_calls(tool_calls), // LlmStreamChunk::PendingToolCall(pending_tool_call) => { // self.process_pending_tool_call(pending_tool_call) @@ -131,11 +146,11 @@ impl LlmStreamWriter { None => break, } - if self.should_flush(&last_flush_time) { + if self.should_flush(&last_flushed_at) { if let Err(err) = self.flush_chunks(writer).await { self.process_error(LlmStreamChunkError::from(err)); } - last_flush_time = Instant::now(); + last_flushed_at = Instant::now(); } } @@ -194,7 +209,7 @@ impl LlmStreamWriter { self.errors.get_or_insert_default().push(err); } - fn should_flush(&self, last_flush_time: &Instant) -> bool { + fn should_flush(&self, last_flushed_at: &Instant) -> bool { // if self.current_chunk.tool_calls.is_some() || self.current_chunk.error.is_some() { // return true; // } @@ -202,7 +217,7 @@ impl LlmStreamWriter { return true; } let text = self.current_chunk.text.as_ref(); - last_flush_time.elapsed() > FLUSH_INTERVAL || text.is_some_and(|t| t.len() > MAX_CHUNK_SIZE) + last_flushed_at.elapsed() > FLUSH_INTERVAL || text.is_some_and(|t| t.len() > MAX_CHUNK_SIZE) } /// Flushes the current chunk(s) to the Redis stream. From c02711c45ab16552833067175548bee210608df8 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Thu, 16 Jul 2026 13:46:34 -0400 Subject: [PATCH 101/111] calculate ttft correctly for prompts --- server-new/src/services/chat/mod.rs | 14 ++++++-------- 1 file changed, 6 insertions(+), 8 deletions(-) diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index 0dfe10a..23007c0 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -112,6 +112,7 @@ impl<'r> ChatService<'r> { }; let log = db.logs().create(create_log).await?; + let start_time = Instant::now(); let request = LlmChatRequest { messages: &[LlmMessage::User(prompt)], options: &options, @@ -132,7 +133,6 @@ impl<'r> ChatService<'r> { // Spawn thread to process LLM streaming response let stream_key = StreamingService::prompt_key(&user_id); - let start_time = Instant::now(); let (stream_access, ws_writer, ws_reader) = StreamingService::new(self.tinistream) .create_stream(&stream_key) .await?; @@ -279,13 +279,11 @@ impl<'r> ChatService<'r> { let log = db.logs().create(create_log).await?; let start_time = Instant::now(); - let (response_stream, meta) = match provider - .stream_chat(LlmChatRequest { - messages: &messages::build_llm_messages(messages)?, - options: ¶ms.chat_options, - }) - .await - { + let request = LlmChatRequest { + messages: &messages::build_llm_messages(messages)?, + options: ¶ms.chat_options, + }; + let (response_stream, meta) = match provider.stream_chat(request).await { Ok(response) => response, Err(err) => { let complete_log = LlmLogComplete { From e57df97a32e5147430a5a72801bbdae7acc2feb1 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Thu, 16 Jul 2026 14:53:58 -0400 Subject: [PATCH 102/111] re-arrange log metadata --- server-new/src/db/models/log.rs | 12 ++++++++++-- server-new/src/db/repositories/log.rs | 13 +++++++++---- 2 files changed, 19 insertions(+), 6 deletions(-) diff --git a/server-new/src/db/models/log.rs b/server-new/src/db/models/log.rs index 50c1d85..7237595 100644 --- a/server-new/src/db/models/log.rs +++ b/server-new/src/db/models/log.rs @@ -76,12 +76,20 @@ pub enum ChatRsLogStatus { #[derive(Debug, Default, Serialize, Deserialize, FromSqlRow, AsExpression, AsJsonb, JsonSchema)] #[diesel(sql_type = diesel::sql_types::Jsonb)] pub struct ChatRsLogMeta { - pub temperature: Option, - pub max_tokens: Option, + /// Options passed to the LLM provider + pub options: Option, + /// Any errors received from the LLM provider pub errors: Option>, + /// The request ID at the LLM provider pub request_id: Option, } +#[derive(Debug, Default, Serialize, Deserialize, JsonSchema)] +pub struct ChatRsLogMetaOptions { + pub temperature: Option, + pub max_tokens: Option, +} + #[derive(Insertable)] #[diesel(table_name = super::schema::llm_logs)] pub struct NewChatRsLog<'a> { diff --git a/server-new/src/db/repositories/log.rs b/server-new/src/db/repositories/log.rs index 103a4d9..9afdd91 100644 --- a/server-new/src/db/repositories/log.rs +++ b/server-new/src/db/repositories/log.rs @@ -8,7 +8,10 @@ use uuid::Uuid; use crate::{ db::{ DbConnection, UtcDateTime, - models::{ChatRsLogKind, ChatRsLogMeta, ChatRsLogStatus, NewChatRsLog, UpdateChatRsLog}, + models::{ + ChatRsLogKind, ChatRsLogMeta, ChatRsLogMetaOptions, ChatRsLogStatus, NewChatRsLog, + UpdateChatRsLog, + }, schema::llm_logs, }, llm::types::{LlmChatOptions, LlmUsage}, @@ -44,8 +47,10 @@ impl<'a> LogRepository<'a> { .unwrap_or_default(), status: ChatRsLogStatus::Started.as_ref(), meta: Some(&ChatRsLogMeta { - temperature: llm_options.and_then(|o| o.temperature), - max_tokens: llm_options.and_then(|o| o.max_tokens), + options: Some(ChatRsLogMetaOptions { + temperature: llm_options.and_then(|o| o.temperature), + max_tokens: llm_options.and_then(|o| o.max_tokens), + }), ..Default::default() }), started_at: chrono::Utc::now(), @@ -82,8 +87,8 @@ impl<'a> LogRepository<'a> { completed_at: Some(completed_at.unwrap_or_else(chrono::Utc::now)), ttft_ms: first_token_in.and_then(|d| d.as_millis().try_into().ok()), meta: ChatRsLogMeta { - request_id: request_id.map(str::to_owned), errors, + request_id: request_id.map(str::to_owned), ..log.meta }, }; From fc71f5f14cc95725d8845882933428e76da6fd94 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Thu, 16 Jul 2026 15:45:51 -0400 Subject: [PATCH 103/111] migrate llm logs from message metadata to table --- .../up.sql | 48 +++++++++++++++++++ server-new/src/api/session.rs | 1 + server-new/src/db/models/chat.rs | 21 +------- server-new/src/db/models/log.rs | 1 + server-new/src/services/chat/mod.rs | 11 +---- 5 files changed, 53 insertions(+), 29 deletions(-) diff --git a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql index 850a878..b47a657 100644 --- a/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql +++ b/server-new/migrations/2026-07-15-053539-0000_add_request_logs/up.sql @@ -1,3 +1,4 @@ +-- Create LLM request logs table CREATE TABLE llm_logs ( id SERIAL PRIMARY KEY, kind text NOT NULL, -- chat, title, prompt, image, audio, etc. @@ -25,3 +26,50 @@ WHERE message_id IS NOT NULL; CREATE INDEX llm_logs_provider_id_started_at_idx ON llm_logs (provider_id, started_at DESC); + +-- Migrate assistant metadata to LLM request logs table +INSERT INTO + llm_logs ( + kind, + user_id, + session_id, + message_id, + provider_id, + model, + input_tokens, + output_tokens, + cost, + status, + meta, + started_at + ) +SELECT + 'chat', + chat_sessions.user_id, + session_id, + chat_messages.id, + providers.id, + coalesce((chat_messages.meta #>> '{assistant,provider_options,model}')::text, ''), + (chat_messages.meta #>> '{assistant,usage,input_tokens}')::int4, + (chat_messages.meta #>> '{assistant,usage,output_tokens}')::int4, + (chat_messages.meta #>> '{assistant,usage,cost}')::numeric(12, 6), + CASE + WHEN (chat_messages.meta #>> '{assistant,partial}')::bool THEN 'cancelled' + WHEN chat_messages.meta @? '$.assistant.errors[0]' THEN 'error' + ELSE 'completed' + END, + jsonb_strip_nulls( + jsonb_build_object( + 'options', + chat_messages.meta #> '{assistant,provider_options}', + 'errors', + chat_messages.meta #> '{assistant,errors}' + ) + ), + chat_messages.created_at +FROM + chat_messages + JOIN chat_sessions ON chat_sessions.id = chat_messages.session_id + LEFT JOIN providers ON providers.id = (chat_messages.meta #>> '{assistant,provider_id}')::int4 +WHERE + chat_messages.role = 'assistant'; diff --git a/server-new/src/api/session.rs b/server-new/src/api/session.rs index 0b1eff2..731ba2b 100644 --- a/server-new/src/api/session.rs +++ b/server-new/src/api/session.rs @@ -164,6 +164,7 @@ struct SessionMessage { /// The message message: ChatRsMessage, /// Request metadata for assistant responses + #[serde(skip_serializing_if = "Option::is_none")] llm_request: Option, } diff --git a/server-new/src/db/models/chat.rs b/server-new/src/db/models/chat.rs index 8224462..d7609ba 100644 --- a/server-new/src/db/models/chat.rs +++ b/server-new/src/db/models/chat.rs @@ -5,10 +5,7 @@ use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use uuid::Uuid; -use crate::{ - db::models::ChatRsUser, - llm::types::{LlmChatOptions, LlmUsage}, -}; +use crate::db::models::ChatRsUser; #[derive(Identifiable, Associations, Queryable, Selectable, Serialize, JsonSchema)] #[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] @@ -110,26 +107,12 @@ pub struct UserMeta { #[derive(Debug, Default, Serialize, Deserialize, JsonSchema)] pub struct AssistantMeta { - /// The ID of the LLM provider used to generate this message - pub provider_id: i32, - /// Options passed to the LLM provider - #[serde(skip_serializing_if = "Option::is_none")] - pub provider_options: Option, - /// The tool calls requested by the assistant + // /// The tool calls requested by the assistant // #[serde(skip_serializing_if = "Option::is_none")] // pub tool_calls: Option>, /// IDs of generated files #[serde(skip_serializing_if = "Option::is_none")] pub files: Option>, - /// Provider usage information - #[serde(skip_serializing_if = "Option::is_none")] - pub usage: Option, - /// Errors encountered during message generation - #[serde(skip_serializing_if = "Option::is_none")] - pub errors: Option>, - /// Whether this is a partial and/or interrupted message - #[serde(skip_serializing_if = "Option::is_none")] - pub partial: Option, } #[derive(Insertable)] diff --git a/server-new/src/db/models/log.rs b/server-new/src/db/models/log.rs index 7237595..01c227a 100644 --- a/server-new/src/db/models/log.rs +++ b/server-new/src/db/models/log.rs @@ -84,6 +84,7 @@ pub struct ChatRsLogMeta { pub request_id: Option, } +#[skip_serializing_none] #[derive(Debug, Default, Serialize, Deserialize, JsonSchema)] pub struct ChatRsLogMetaOptions { pub temperature: Option, diff --git a/server-new/src/services/chat/mod.rs b/server-new/src/services/chat/mod.rs index 23007c0..3d370bc 100644 --- a/server-new/src/services/chat/mod.rs +++ b/server-new/src/services/chat/mod.rs @@ -332,16 +332,7 @@ impl<'r> ChatService<'r> { let completed_at = chrono::Utc::now(); let mut db = DbService::from_pool(&db_pool).await?; - let assistant_meta = AssistantMeta { - provider_id: params.provider_id, - provider_options: Some(params.chat_options), - // tool_calls: response.tool_calls, - // files: image_ids, - usage: output.usage, - errors: output.errors.clone(), - partial: output.cancelled.then_some(true), - ..Default::default() - }; + let assistant_meta = AssistantMeta::default(); let new_message = db .chats() .save_message(NewChatRsMessage { From 3eb284e476ce68ffc99443b9178440aa1ab9f428 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Fri, 17 Jul 2026 00:29:15 -0400 Subject: [PATCH 104/111] checkpoint: file storage --- server-new/Cargo.lock | 33 +++++++ server-new/Cargo.toml | 2 +- server-new/src/config.rs | 4 +- server-new/src/services/mod.rs | 1 + server-new/src/services/storage/engines.rs | 3 + .../src/services/storage/engines/local.rs | 91 +++++++++++++++++++ server-new/src/services/storage/error.rs | 9 ++ server-new/src/services/storage/interface.rs | 18 ++++ server-new/src/services/storage/mod.rs | 50 ++++++++++ 9 files changed, 209 insertions(+), 2 deletions(-) create mode 100644 server-new/src/services/storage/engines.rs create mode 100644 server-new/src/services/storage/engines/local.rs create mode 100644 server-new/src/services/storage/error.rs create mode 100644 server-new/src/services/storage/interface.rs create mode 100644 server-new/src/services/storage/mod.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index 2591421..f5a486d 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -249,6 +249,7 @@ dependencies = [ "matchit", "memchr", "mime", + "multer", "percent-encoding", "pin-project-lite", "serde_core", @@ -937,6 +938,15 @@ version = "1.16.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" +[[package]] +name = "encoding_rs" +version = "0.8.35" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75030f3c4f45dafd7586dd6780965a8c7e8e285a5ecb86713e63a79c5b2766f3" +dependencies = [ + "cfg-if", +] + [[package]] name = "equivalent" version = "1.0.2" @@ -1774,6 +1784,23 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "multer" +version = "3.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "83e87776546dc87511aa5ee218730c92b666d7264ab6ed41f9d215af9cd5224b" +dependencies = [ + "bytes", + "encoding_rs", + "futures-util", + "http", + "httparse", + "memchr", + "mime", + "spin", + "version_check", +] + [[package]] name = "nom" version = "7.1.3" @@ -2914,6 +2941,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "spin" +version = "0.9.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e" + [[package]] name = "stable_deref_trait" version = "1.2.1" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 91ab789..1d9adad 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -15,7 +15,7 @@ aide = { anyhow = "1.0.103" async-stream = "0.3.6" async-trait = "0.1.89" -axum = { version = "0.8.9", features = ["json", "query"] } +axum = { version = "0.8.9", features = ["json", "multipart", "query"] } axum-aide-macros = { git = "https://git.fasharp.io/fa-sharp/axum-aide-macros", rev = "5b00e645df" diff --git a/server-new/src/config.rs b/server-new/src/config.rs index af304a4..e922b14 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -1,4 +1,4 @@ -use std::net::IpAddr; +use std::{net::IpAddr, path::PathBuf}; use axum_plugin::figment::{ Figment, @@ -36,6 +36,7 @@ pub struct ServerConfig { pub port: u16, pub base_url: String, pub log_level: String, + pub data_dir: PathBuf, pub web_root: String, pub request_id_header: String, pub ip_header: Option, @@ -47,6 +48,7 @@ impl Default for ServerConfig { port: 8080, base_url: String::from("http://localhost:8080"), log_level: String::from("info"), + data_dir: PathBuf::from("/data"), web_root: String::from("../web/dist"), request_id_header: String::from("x-request-id"), ip_header: None, diff --git a/server-new/src/services/mod.rs b/server-new/src/services/mod.rs index 95a51c8..b68e8b8 100644 --- a/server-new/src/services/mod.rs +++ b/server-new/src/services/mod.rs @@ -2,4 +2,5 @@ pub mod auth; pub mod chat; pub mod model; pub mod provider; +pub mod storage; pub mod stream; diff --git a/server-new/src/services/storage/engines.rs b/server-new/src/services/storage/engines.rs new file mode 100644 index 0000000..4b15657 --- /dev/null +++ b/server-new/src/services/storage/engines.rs @@ -0,0 +1,3 @@ +mod local; + +pub use local::LocalStorage; diff --git a/server-new/src/services/storage/engines/local.rs b/server-new/src/services/storage/engines/local.rs new file mode 100644 index 0000000..805b214 --- /dev/null +++ b/server-new/src/services/storage/engines/local.rs @@ -0,0 +1,91 @@ +use std::path::{Path, PathBuf}; + +use futures::future::BoxFuture; +use tokio::{ + fs::File, + io::{AsyncRead, AsyncReadExt, AsyncWriteExt, BufWriter}, +}; + +use crate::services::storage::{ + StorageEngine, + error::{StorageError, StorageResult}, +}; + +/// Default storage engine using local filesystem +pub struct LocalStorage { + base_path: PathBuf, +} + +impl LocalStorage { + pub fn new(base_path: &str) -> Self { + Self { + base_path: PathBuf::from(base_path), + } + } + + fn file_path(&self, path: &Path) -> PathBuf { + self.base_path.join(path) + } + + async fn file_exists(path: &Path) -> bool { + match tokio::fs::metadata(&path).await { + Ok(meta) => meta.is_file(), + Err(_) => false, + } + } +} + +impl StorageEngine for LocalStorage { + fn create<'r>( + &self, + path: &Path, + reader: &'r mut (dyn AsyncRead + Unpin + Send), + ) -> BoxFuture<'r, StorageResult> { + let file_path = self.file_path(path); + + Box::pin(async move { + let dir = file_path.parent().expect("should have a parent directory"); + tokio::fs::create_dir_all(&dir).await?; + + let mut file = File::create_new(&file_path).await?; + let mut file_writer = BufWriter::new(&mut file); + let mut read_buffer = [0; 4096]; + let mut total_bytes_written: usize = 0; + + loop { + let n = reader.read(&mut read_buffer).await?; + if n == 0 { + break; + } + file_writer.write_all(&read_buffer[..n]).await?; + total_bytes_written += n; + } + + file_writer.flush().await?; + file.sync_all().await?; + + Ok(total_bytes_written) + }) + } + + fn exists(&self, path: &Path) -> BoxFuture<'_, StorageResult> { + let file_path = self.file_path(path); + + Box::pin(async move { Ok(Self::file_exists(&file_path).await) }) + } + + fn delete(&self, path: &Path) -> BoxFuture<'_, StorageResult<()>> { + let file_path = self.file_path(path); + + Box::pin(async move { + match Self::file_exists(&file_path).await { + true => Ok(tokio::fs::remove_file(&file_path).await?), + false => Err(StorageError::NotFound), + } + }) + } + + fn signed_url(&self, path: &Path) -> StorageResult { + todo!() + } +} diff --git a/server-new/src/services/storage/error.rs b/server-new/src/services/storage/error.rs new file mode 100644 index 0000000..1afcb91 --- /dev/null +++ b/server-new/src/services/storage/error.rs @@ -0,0 +1,9 @@ +pub type StorageResult = Result; + +#[derive(Debug, thiserror::Error)] +pub enum StorageError { + #[error("IO error: {0}")] + Io(#[from] std::io::Error), + #[error("File not found")] + NotFound, +} diff --git a/server-new/src/services/storage/interface.rs b/server-new/src/services/storage/interface.rs new file mode 100644 index 0000000..96ebe69 --- /dev/null +++ b/server-new/src/services/storage/interface.rs @@ -0,0 +1,18 @@ +use std::path::Path; + +use futures::future::BoxFuture; +use tokio::io::AsyncRead; + +use super::error::StorageResult; + +/// Trait representing an underlying storage to manage files for LLM chats and responses +pub trait StorageEngine: Send + Sync { + fn create<'r>( + &self, + path: &Path, + reader: &'r mut (dyn AsyncRead + Unpin + Send), + ) -> BoxFuture<'r, StorageResult>; + fn exists(&self, path: &Path) -> BoxFuture<'_, StorageResult>; + fn delete(&self, path: &Path) -> BoxFuture<'_, StorageResult<()>>; + fn signed_url(&self, path: &Path) -> StorageResult; +} diff --git a/server-new/src/services/storage/mod.rs b/server-new/src/services/storage/mod.rs new file mode 100644 index 0000000..0c2c0b3 --- /dev/null +++ b/server-new/src/services/storage/mod.rs @@ -0,0 +1,50 @@ +use std::path::PathBuf; + +use axum::extract::multipart::MultipartError; +use futures::{Stream, TryStreamExt}; +use tokio_util::io::StreamReader; +use uuid::Uuid; + +use crate::services::storage::{engines::LocalStorage, error::StorageResult}; + +mod engines; +mod error; +mod interface; + +pub use interface::StorageEngine; + +pub struct StorageService<'r> { + data_dir: &'r str, +} + +impl<'r> StorageService<'r> { + fn file_path(&self, user_id: &Uuid, session_id: &Uuid, name: &str) -> PathBuf { + let session_folder = PathBuf::from(format!("{user_id}/{session_id}")); + + session_folder.join(name) + } + + pub async fn create_file( + &self, + user_id: &Uuid, + session_id: &Uuid, + name: &str, + stream: impl Stream> + Send + Unpin, + ) -> StorageResult<()> { + let storage: Box = + Box::new(LocalStorage::new(&format!("{}/storage", self.data_dir))); + + let path = self.file_path(user_id, session_id, name); + let mut reader = StreamReader::new(stream.map_err(std::io::Error::other)); + + let size = match storage.create(&path, &mut reader).await { + Ok(n) => n, + Err(err) => { + let _ = storage.delete(&path).await; + return Err(err); + } + }; + + todo!() + } +} From 81b87bac796c28f8c4ad21f127301ddbca7d0464 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Fri, 17 Jul 2026 10:23:58 -0400 Subject: [PATCH 105/111] checkpoint: file database models and repository --- .../down.sql | 1 + .../up.sql | 6 ++ server-new/src/db/mod.rs | 3 + server-new/src/db/models.rs | 4 +- server-new/src/db/models/file.rs | 65 ++++++++++++++ server-new/src/db/repositories.rs | 2 + server-new/src/db/repositories/chat.rs | 2 +- server-new/src/db/repositories/file.rs | 90 +++++++++++++++++++ server-new/src/db/schema.rs | 10 +++ 9 files changed, 180 insertions(+), 3 deletions(-) create mode 100644 server-new/migrations/2026-07-17-043348-0000_add_message_attachments/down.sql create mode 100644 server-new/migrations/2026-07-17-043348-0000_add_message_attachments/up.sql create mode 100644 server-new/src/db/models/file.rs create mode 100644 server-new/src/db/repositories/file.rs diff --git a/server-new/migrations/2026-07-17-043348-0000_add_message_attachments/down.sql b/server-new/migrations/2026-07-17-043348-0000_add_message_attachments/down.sql new file mode 100644 index 0000000..75cb30b --- /dev/null +++ b/server-new/migrations/2026-07-17-043348-0000_add_message_attachments/down.sql @@ -0,0 +1 @@ +DROP TABLE message_attachments; diff --git a/server-new/migrations/2026-07-17-043348-0000_add_message_attachments/up.sql b/server-new/migrations/2026-07-17-043348-0000_add_message_attachments/up.sql new file mode 100644 index 0000000..bc170bc --- /dev/null +++ b/server-new/migrations/2026-07-17-043348-0000_add_message_attachments/up.sql @@ -0,0 +1,6 @@ +-- Create table tracking file attachments to messaages +CREATE TABLE message_attachments ( + message_id uuid REFERENCES chat_messages (id) ON DELETE CASCADE, + file_id uuid REFERENCES files (id) ON DELETE CASCADE, + PRIMARY KEY (message_id, file_id) +); diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs index a448e87..b88ce36 100644 --- a/server-new/src/db/mod.rs +++ b/server-new/src/db/mod.rs @@ -59,6 +59,9 @@ impl DbService { pub fn chats(&mut self) -> repositories::ChatRepository<'_> { repositories::ChatRepository::new(&mut self.cxn) } + pub fn files(&mut self) -> repositories::FileRepository<'_> { + repositories::FileRepository::new(&mut self.cxn) + } pub fn logs(&mut self) -> repositories::LogRepository<'_> { repositories::LogRepository::new(&mut self.cxn) } diff --git a/server-new/src/db/models.rs b/server-new/src/db/models.rs index 38ce052..ac65177 100644 --- a/server-new/src/db/models.rs +++ b/server-new/src/db/models.rs @@ -4,7 +4,7 @@ use crate::db::schema; mod api_key; mod chat; -// mod file; +mod file; mod log; mod provider; mod secret; @@ -14,7 +14,7 @@ mod user; pub use api_key::*; pub use chat::*; -// pub use file::*; +pub use file::*; pub use log::*; pub use provider::*; pub use secret::*; diff --git a/server-new/src/db/models/file.rs b/server-new/src/db/models/file.rs new file mode 100644 index 0000000..98a1ed7 --- /dev/null +++ b/server-new/src/db/models/file.rs @@ -0,0 +1,65 @@ +use diesel::prelude::*; +use schemars::JsonSchema; +use strum::{AsRefStr, EnumString}; +use uuid::Uuid; + +use crate::db::{ + UtcDateTime, + models::{ChatRsMessage, ChatRsUser}, +}; + +#[derive(Identifiable, Associations, Queryable, Selectable, JsonSchema, serde::Serialize)] +#[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] +#[diesel(table_name = super::schema::files)] +pub struct ChatRsFile { + pub id: Uuid, + #[serde(skip)] + pub user_id: Uuid, + pub session_id: Option, + pub path: String, + #[schemars(with = "ChatRsFileType")] + pub file_type: String, + pub content_type: String, + pub size: i32, + pub created_at: UtcDateTime, + #[serde(skip)] + pub updated_at: UtcDateTime, +} + +#[derive(Identifiable, Selectable, Queryable, Associations)] +#[diesel(belongs_to(ChatRsMessage, foreign_key = message_id))] +#[diesel(belongs_to(ChatRsFile, foreign_key = file_id))] +#[diesel(table_name = super::schema::message_attachments)] +#[diesel(primary_key(message_id, file_id))] +pub struct ChatRsMessageAttachment { + pub message_id: Uuid, + pub file_id: Uuid, +} + +#[derive(Insertable)] +#[diesel(table_name = super::schema::files)] +pub struct NewChatRsFile<'r> { + pub user_id: &'r Uuid, + pub session_id: Option<&'r Uuid>, + pub path: &'r str, + pub file_type: &'r str, + pub content_type: &'r str, + pub size: i32, +} + +#[derive(Insertable)] +#[diesel(table_name = super::schema::message_attachments)] +pub struct NewChatRsMessageAttachment<'r> { + pub message_id: &'r Uuid, + pub file_id: &'r Uuid, +} + +/// File modality +#[derive(Debug, PartialEq, Eq, Hash, EnumString, AsRefStr, JsonSchema)] +#[serde(rename_all = "lowercase")] +#[strum(serialize_all = "lowercase")] +pub enum ChatRsFileType { + Text, + Image, + Pdf, +} diff --git a/server-new/src/db/repositories.rs b/server-new/src/db/repositories.rs index 31aac27..2698803 100644 --- a/server-new/src/db/repositories.rs +++ b/server-new/src/db/repositories.rs @@ -2,6 +2,7 @@ mod api_key; mod chat; +mod file; mod log; mod provider; mod secret; @@ -10,6 +11,7 @@ mod user; pub use api_key::ApiKeyRepository; pub use chat::ChatRepository; +pub use file::FileRepository; pub use log::{LlmLogComplete, LlmLogCreate, LogRepository}; pub use provider::ProviderRepository; pub use secret::SecretRepository; diff --git a/server-new/src/db/repositories/chat.rs b/server-new/src/db/repositories/chat.rs index f9d835b..a3d2798 100644 --- a/server-new/src/db/repositories/chat.rs +++ b/server-new/src/db/repositories/chat.rs @@ -63,7 +63,7 @@ impl<'a> ChatRepository<'a> { message_id: &Uuid, ) -> Result, diesel::result::Error> { chat_messages::table - .inner_join(chat_sessions::table.on(chat_sessions::id.eq(chat_messages::session_id))) + .inner_join(chat_sessions::table) .select(ChatRsMessage::as_select()) .filter(chat_sessions::user_id.eq(user_id)) .filter(chat_messages::id.eq(message_id)) diff --git a/server-new/src/db/repositories/file.rs b/server-new/src/db/repositories/file.rs new file mode 100644 index 0000000..99b6306 --- /dev/null +++ b/server-new/src/db/repositories/file.rs @@ -0,0 +1,90 @@ +use diesel::prelude::*; +use diesel_async::RunQueryDsl; +use uuid::Uuid; + +use crate::db::{ + DbConnection, + models::{ChatRsFile, ChatRsMessageAttachment, NewChatRsFile, NewChatRsMessageAttachment}, + schema::{files, message_attachments}, +}; + +pub struct FileRepository<'a> { + pub db: &'a mut DbConnection, +} + +impl<'a> FileRepository<'a> { + pub fn new(db: &'a mut DbConnection) -> Self { + Self { db } + } + + pub async fn create_session_file( + &mut self, + file: NewChatRsFile<'_>, + ) -> QueryResult { + diesel::insert_into(files::table) + .values(file) + .returning(ChatRsFile::as_returning()) + .get_result(self.db) + .await + } + + pub async fn find_session_file( + &mut self, + user_id: &Uuid, + session_id: &Uuid, + file_id: &Uuid, + ) -> QueryResult { + files::table + .filter(files::user_id.eq(user_id)) + .filter(files::session_id.eq(session_id)) + .filter(files::id.eq(file_id)) + .select(ChatRsFile::as_select()) + .first(self.db) + .await + } + + pub async fn list_session_files( + &mut self, + user_id: &Uuid, + session_id: &Uuid, + ) -> QueryResult> { + files::table + .filter(files::user_id.eq(user_id)) + .filter(files::session_id.eq(session_id)) + .select(ChatRsFile::as_select()) + .load(self.db) + .await + } + + pub async fn attach_files( + &mut self, + message_id: &Uuid, + file_ids: &[Uuid], + ) -> QueryResult { + let attachments = file_ids.iter().map(|file_id| NewChatRsMessageAttachment { + message_id, + file_id, + }); + + diesel::insert_into(message_attachments::table) + .values(attachments.collect::>()) + .returning(ChatRsMessageAttachment::as_returning()) + .get_result(self.db) + .await + } + + pub async fn delete_session_file( + &mut self, + user_id: &Uuid, + session_id: &Uuid, + file_id: &Uuid, + ) -> QueryResult { + diesel::delete(files::table) + .filter(files::user_id.eq(user_id)) + .filter(files::session_id.eq(session_id)) + .filter(files::id.eq(file_id)) + .returning(files::id) + .get_result(self.db) + .await + } +} diff --git a/server-new/src/db/schema.rs b/server-new/src/db/schema.rs index 07aff3b..c3cf047 100644 --- a/server-new/src/db/schema.rs +++ b/server-new/src/db/schema.rs @@ -104,6 +104,13 @@ diesel::table! { } } +diesel::table! { + message_attachments (message_id, file_id) { + message_id -> Uuid, + file_id -> Uuid, + } +} + diesel::table! { providers (id) { id -> Int4, @@ -165,6 +172,8 @@ diesel::joinable!(llm_logs -> chat_messages (message_id)); diesel::joinable!(llm_logs -> chat_sessions (session_id)); diesel::joinable!(llm_logs -> providers (provider_id)); diesel::joinable!(llm_logs -> users (user_id)); +diesel::joinable!(message_attachments -> chat_messages (message_id)); +diesel::joinable!(message_attachments -> files (file_id)); diesel::joinable!(providers -> secrets (api_key_id)); diesel::joinable!(providers -> users (user_id)); diesel::joinable!(secrets -> users (user_id)); @@ -178,6 +187,7 @@ diesel::allow_tables_to_appear_in_same_query!( external_api_tools, files, llm_logs, + message_attachments, providers, secrets, system_tools, From 521729f46d87ac6436925f774f95fbf65beb6ca2 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sat, 18 Jul 2026 15:00:52 -0400 Subject: [PATCH 106/111] db query logging --- server-new/src/plugins/database.rs | 18 +++++++++++++++++- 1 file changed, 17 insertions(+), 1 deletion(-) diff --git a/server-new/src/plugins/database.rs b/server-new/src/plugins/database.rs index bd92ab8..94ca026 100644 --- a/server-new/src/plugins/database.rs +++ b/server-new/src/plugins/database.rs @@ -1,9 +1,11 @@ use anyhow::Context; +use diesel::connection::InstrumentationEvent; use diesel_async::{ - AsyncMigrationHarness, AsyncPgConnection, + AsyncConnection, AsyncMigrationHarness, AsyncPgConnection, pooled_connection::{AsyncDieselConnectionManager, ManagerConfig, deadpool::Pool}, }; use diesel_migrations::{EmbeddedMigrations, MigrationHarness}; +use futures::TryFutureExt; use crate::{db::DbPool, plugins::AxumPlugin}; @@ -18,6 +20,20 @@ pub fn plugin() -> AxumPlugin { let mut config = ManagerConfig::default(); config.recycling_method = diesel_async::pooled_connection::RecyclingMethod::Fast; + config.custom_setup = Box::new(|url| { + Box::pin(AsyncPgConnection::establish(url).map_ok(|mut conn| { + conn.set_instrumentation(|ev: InstrumentationEvent<'_>| { + if let InstrumentationEvent::FinishQuery { query, error, .. } = ev { + if let Some(err) = error { + tracing::error!(?query, ?err, "Failed to execute query"); + } else { + tracing::debug!(?query); + } + }; + }); + conn + })) + }); config }, ); From 2b8f7aad3ece0f4b5198d7cc869430dc8c8c0b35 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sat, 18 Jul 2026 15:24:25 -0400 Subject: [PATCH 107/111] storage routes --- server-new/.gitignore | 3 + server-new/Cargo.toml | 9 +- server-new/config.toml | 1 + server-new/src/api/chat.rs | 10 +- server-new/src/api/mod.rs | 26 ++- server-new/src/api/session.rs | 2 +- server-new/src/api/storage.rs | 86 ++++++++++ server-new/src/api/upload.rs | 78 +++++++++ server-new/src/config.rs | 6 +- server-new/src/db/mod.rs | 7 +- server-new/src/db/models/file.rs | 5 +- server-new/src/db/repositories/chat.rs | 2 +- server-new/src/db/repositories/file.rs | 71 ++++++-- server-new/src/services/storage/engines.rs | 2 + .../src/services/storage/engines/local.rs | 59 +++---- server-new/src/services/storage/error.rs | 34 +++- server-new/src/services/storage/interface.rs | 14 +- server-new/src/services/storage/mod.rs | 152 +++++++++++++++--- server-new/src/state.rs | 4 + 19 files changed, 473 insertions(+), 98 deletions(-) create mode 100644 server-new/src/api/storage.rs create mode 100644 server-new/src/api/upload.rs diff --git a/server-new/.gitignore b/server-new/.gitignore index 1588a67..a0c632a 100644 --- a/server-new/.gitignore +++ b/server-new/.gitignore @@ -1,3 +1,6 @@ +# Local storage files +.local/ + # Rust /target .diesel_lock diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index 1d9adad..d5ae512 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -10,7 +10,14 @@ aes-gcm = "0.11.0" aide = { git = "https://github.com/hniksic/aide.git", rev = "7246c20", - features = ["axum", "axum-json", "axum-query", "macros", "swagger"] + features = [ + "axum", + "axum-json", + "axum-multipart", + "axum-query", + "macros", + "swagger" + ] } anyhow = "1.0.103" async-stream = "0.3.6" diff --git a/server-new/config.toml b/server-new/config.toml index c5bc5ae..289f384 100644 --- a/server-new/config.toml +++ b/server-new/config.toml @@ -5,6 +5,7 @@ host = "127.0.0.1" port = 8080 base_url = "http://localhost:8080" log_level = "info" +data_dir = ".local" [database] url = "postgres://postgres:postgres@localhost/postgres" diff --git a/server-new/src/api/chat.rs b/server-new/src/api/chat.rs index f864286..8d21cd5 100644 --- a/server-new/src/api/chat.rs +++ b/server-new/src/api/chat.rs @@ -19,11 +19,11 @@ api_routes! { state: AppState, tag: ApiTag::Chat.into(), POST "/prompt" => prompt, "Prompt"; - GET "/sessions" => get_active_streams, "Get active chat streams"; - GET "/sessions/{session_id}" => connect_chat_stream, "Access active chat stream"; - POST "/sessions/{session_id}" => chat_stream, "Stream chat session response"; - POST "/sessions/{session_id}/cancel" => cancel_chat_stream, "Cancel active chat stream"; - POST "/sessions/{session_id}/regenerate" => regenerate_response, "Regenerate chat response"; + GET "/session" => get_active_streams, "Get sessions with active streams"; + GET "/session/{session_id}" => connect_chat_stream, "Access active chat stream"; + POST "/session/{session_id}" => chat_stream, "Stream chat session response"; + POST "/session/{session_id}/cancel" => cancel_chat_stream, "Cancel active chat stream"; + POST "/session/{session_id}/regenerate" => regenerate_response, "Regenerate chat response"; } async fn get_active_streams( diff --git a/server-new/src/api/mod.rs b/server-new/src/api/mod.rs index 2ac879a..c0d4f56 100644 --- a/server-new/src/api/mod.rs +++ b/server-new/src/api/mod.rs @@ -5,23 +5,26 @@ use aide::{ openapi::{OpenApi, SecurityScheme, Server}, swagger::Swagger, }; -use axum::{Extension, routing::get}; +use axum::{Extension, extract::DefaultBodyLimit, routing::get}; use axum_plugin::AdHocPlugin; use strum::{Display, EnumIter, EnumMessage, IntoEnumIterator, IntoStaticStr}; use crate::{config::AppConfig, state::AppState}; -pub mod api_key; -pub mod auth; -pub mod chat; -pub mod health; -pub mod provider; -pub mod session; +mod api_key; +mod auth; +mod chat; +mod health; +mod provider; +mod session; +mod storage; +mod upload; const API_BASE: &str = "/api/v1"; const API_AUTH_BASE: &str = "/api/v1/auth"; pub const API_KEY_SCHEME: &str = "ApiKey"; +/// API route tags for OpenAPI docs #[derive(Display, IntoStaticStr, EnumMessage, EnumIter)] enum ApiTag { #[strum(message = "Manage API keys")] @@ -32,11 +35,13 @@ enum ApiTag { Chat, #[strum(message = "AI / LLM Providers")] Provider, + #[strum(message = "Files and attachments")] + Storage, } /// Adds all API routes with OpenAPI docs to the server under `/api/v1` pub fn plugin() -> AdHocPlugin { - AdHocPlugin::named("API routes").on_setup(|_app, router| { + AdHocPlugin::::named("API routes").on_setup(|app, router| { let mut openapi = OpenApi::default(); let api_routes = ApiRouter::new() .nest("/api_key", api_key::routes()) @@ -48,6 +53,11 @@ pub fn plugin() -> AdHocPlugin { .nest("/health", health::routes()) .nest("/provider", provider::routes()) .nest("/session", session::routes()) + .nest("/storage", storage::routes()) + .nest( + "/upload", + upload::routes().layer(DefaultBodyLimit::max(app.config().security.upload_limit)), + ) .finish_api_with(&mut openapi, build_openapi_doc) .route( "/docs/openapi.json", diff --git a/server-new/src/api/session.rs b/server-new/src/api/session.rs index 731ba2b..d20346a 100644 --- a/server-new/src/api/session.rs +++ b/server-new/src/api/session.rs @@ -40,7 +40,7 @@ async fn get_recent_sessions( CurrentUser { user_id }: CurrentUser, Database(mut db): Database, ) -> AppResult>> { - let sessions = db.chats().get_recent_sessions(&user_id).await?; + let sessions = db.chats().list_recent_sessions(&user_id).await?; Ok(Json(sessions)) } diff --git a/server-new/src/api/storage.rs b/server-new/src/api/storage.rs new file mode 100644 index 0000000..b057a4b --- /dev/null +++ b/server-new/src/api/storage.rs @@ -0,0 +1,86 @@ +use axum::{ + Json, + extract::{Path, State}, +}; +use axum_aide_macros::api_routes; +use schemars::JsonSchema; +use serde::Serialize; +use uuid::Uuid; + +use crate::{ + api::ApiTag, + db::models::{ChatRsFile, ChatRsMessageAttachment}, + error::AppResult, + extractors::{CurrentUser, Database}, + state::AppState, +}; + +api_routes! { + state: AppState, + tag: ApiTag::Storage.into(), + GET "/user" => list_user_files, "List user files"; + DELETE "/user/{file_id}" => delete_user_file, "Delete user file"; + GET "/session/{session_id}" => list_session_files, "List session files"; + DELETE "/session/{session_id}/{file_id}" => delete_session_file, "Delete session file"; +} + +async fn list_user_files( + CurrentUser { user_id }: CurrentUser, + Database(mut db): Database, +) -> AppResult>> { + let files = db.files().list_user_files(&user_id).await?; + + Ok(Json(files)) +} + +async fn list_session_files( + CurrentUser { user_id }: CurrentUser, + Path(session_id): Path, + Database(mut db): Database, +) -> AppResult> { + let (files, attachments) = db + .files() + .list_session_files_and_attachments(&user_id, &session_id) + .await?; + + Ok(Json(SessionFilesAndAttachments { files, attachments })) +} + +async fn delete_user_file( + CurrentUser { user_id }: CurrentUser, + Path(file_id): Path, + Database(mut db): Database, + State(state): State, +) -> AppResult> { + let file_id = state + .storage_service() + .delete_file(&mut db, &user_id, None, &file_id) + .await?; + + Ok(Json(FileIdResponse { file_id })) +} + +async fn delete_session_file( + CurrentUser { user_id }: CurrentUser, + Path((session_id, file_id)): Path<(Uuid, Uuid)>, + Database(mut db): Database, + State(state): State, +) -> AppResult> { + let file_id = state + .storage_service() + .delete_file(&mut db, &user_id, Some(&session_id), &file_id) + .await?; + + Ok(Json(FileIdResponse { file_id })) +} + +#[derive(Serialize, JsonSchema)] +struct SessionFilesAndAttachments { + files: Vec, + attachments: Vec, +} + +#[derive(Serialize, JsonSchema)] +struct FileIdResponse { + file_id: Uuid, +} diff --git a/server-new/src/api/upload.rs b/server-new/src/api/upload.rs new file mode 100644 index 0000000..1560afb --- /dev/null +++ b/server-new/src/api/upload.rs @@ -0,0 +1,78 @@ +use axum::{ + Json, + extract::{Multipart, Path, State}, +}; +use axum_aide_macros::api_routes; +use futures::TryStreamExt; +use uuid::Uuid; + +use crate::{ + api::ApiTag, + db::models::ChatRsFile, + error::{AppError, AppResult}, + extractors::{CurrentUser, Database}, + state::AppState, +}; + +api_routes! { + state: AppState, + tag: ApiTag::Storage.into(), + POST "/user/{*file_path}" => upload_user_file, "Upload user file" { + description: "Upload a file to the user account. The file must be the only field in the form, + with a supported content type." + }; + POST "/session/{session_id}/{*file_path}" => upload_session_file, "Upload session file" { + description: "Upload a file to a chat session. The file must be the only field in the form, + with a supported content type." + }; +} + +async fn upload_user_file( + CurrentUser { user_id }: CurrentUser, + Path(path): Path, + State(state): State, + mut multipart: Multipart, +) -> AppResult> { + let field = multipart + .next_field() + .await + .map_err(|err| AppError::bad_request(err.body_text()))? + .ok_or_else(|| AppError::bad_request("no file in request"))?; + let mime = field.content_type().map(str::to_owned); + let stream = field.map_err(std::io::Error::other); + + let file = state + .storage_service() + .create_file(&user_id, None, &path, mime.as_deref(), stream) + .await?; + + Ok(Json(file)) +} + +async fn upload_session_file( + CurrentUser { user_id }: CurrentUser, + Path((sess_id, path)): Path<(Uuid, String)>, + Database(mut db): Database, + State(state): State, + mut multipart: Multipart, +) -> AppResult> { + if db.chats().find_session(&user_id, &sess_id).await?.is_none() { + return Err(AppError::not_found("session not found")); + } + drop(db); // free database connection since this could be long-running request + + let field = multipart + .next_field() + .await + .map_err(|err| AppError::bad_request(err.body_text()))? + .ok_or_else(|| AppError::bad_request("no file in request"))?; + let mime = field.content_type().map(str::to_owned); + let stream = field.map_err(std::io::Error::other); + + let file = state + .storage_service() + .create_file(&user_id, Some(&sess_id), &path, mime.as_deref(), stream) + .await?; + + Ok(Json(file)) +} diff --git a/server-new/src/config.rs b/server-new/src/config.rs index e922b14..5bcbf03 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -111,13 +111,15 @@ impl Default for AuthConfig { #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SecurityConfig { pub body_limit: usize, + pub upload_limit: usize, pub request_timeout: u64, } impl Default for SecurityConfig { fn default() -> Self { Self { - body_limit: 2097152, // 2 MB - request_timeout: 120, // 2 minutes + body_limit: 5 * 1024 * 1024, // 5 MB + upload_limit: 5 * 1024 * 1024, // 5 MB + request_timeout: 120, // 2 minutes } } } diff --git a/server-new/src/db/mod.rs b/server-new/src/db/mod.rs index b88ce36..f4230c9 100644 --- a/server-new/src/db/mod.rs +++ b/server-new/src/db/mod.rs @@ -18,7 +18,7 @@ pub type DbPool = Pool; pub type DbPoolError = PoolError; /// The database connection retrieved from the pool. For pipelining multiple -/// queries in Diesel, a shared reference can be used with `&mut &**conn`. +/// queries in Diesel, a shared reference can be used with `&mut conn.as_ref()`. pub struct DbConnection(Object); impl Deref for DbConnection { type Target = AsyncPgConnection; @@ -31,6 +31,11 @@ impl DerefMut for DbConnection { &mut self.0 } } +impl AsRef for DbConnection { + fn as_ref(&self) -> &AsyncPgConnection { + &**self + } +} /// Date/time format used in all database tables pub type UtcDateTime = chrono::DateTime; diff --git a/server-new/src/db/models/file.rs b/server-new/src/db/models/file.rs index 98a1ed7..670967c 100644 --- a/server-new/src/db/models/file.rs +++ b/server-new/src/db/models/file.rs @@ -1,5 +1,6 @@ use diesel::prelude::*; use schemars::JsonSchema; +use serde::Serialize; use strum::{AsRefStr, EnumString}; use uuid::Uuid; @@ -8,7 +9,7 @@ use crate::db::{ models::{ChatRsMessage, ChatRsUser}, }; -#[derive(Identifiable, Associations, Queryable, Selectable, JsonSchema, serde::Serialize)] +#[derive(Identifiable, Associations, Queryable, Selectable, Serialize, JsonSchema)] #[diesel(belongs_to(ChatRsUser, foreign_key = user_id))] #[diesel(table_name = super::schema::files)] pub struct ChatRsFile { @@ -26,7 +27,7 @@ pub struct ChatRsFile { pub updated_at: UtcDateTime, } -#[derive(Identifiable, Selectable, Queryable, Associations)] +#[derive(Identifiable, Selectable, Queryable, Associations, Serialize, JsonSchema)] #[diesel(belongs_to(ChatRsMessage, foreign_key = message_id))] #[diesel(belongs_to(ChatRsFile, foreign_key = file_id))] #[diesel(table_name = super::schema::message_attachments)] diff --git a/server-new/src/db/repositories/chat.rs b/server-new/src/db/repositories/chat.rs index a3d2798..55aca3b 100644 --- a/server-new/src/db/repositories/chat.rs +++ b/server-new/src/db/repositories/chat.rs @@ -86,7 +86,7 @@ impl<'a> ChatRepository<'a> { .optional() } - pub async fn get_recent_sessions( + pub async fn list_recent_sessions( &mut self, user_id: &Uuid, ) -> Result, diesel::result::Error> { diff --git a/server-new/src/db/repositories/file.rs b/server-new/src/db/repositories/file.rs index 99b6306..a72fc0d 100644 --- a/server-new/src/db/repositories/file.rs +++ b/server-new/src/db/repositories/file.rs @@ -5,7 +5,7 @@ use uuid::Uuid; use crate::db::{ DbConnection, models::{ChatRsFile, ChatRsMessageAttachment, NewChatRsFile, NewChatRsMessageAttachment}, - schema::{files, message_attachments}, + schema::{chat_messages, chat_sessions, files, message_attachments}, }; pub struct FileRepository<'a> { @@ -17,10 +17,7 @@ impl<'a> FileRepository<'a> { Self { db } } - pub async fn create_session_file( - &mut self, - file: NewChatRsFile<'_>, - ) -> QueryResult { + pub async fn create_file(&mut self, file: NewChatRsFile<'_>) -> QueryResult { diesel::insert_into(files::table) .values(file) .returning(ChatRsFile::as_returning()) @@ -28,32 +25,35 @@ impl<'a> FileRepository<'a> { .await } - pub async fn find_session_file( + pub async fn find_user_file( &mut self, user_id: &Uuid, - session_id: &Uuid, file_id: &Uuid, - ) -> QueryResult { + ) -> QueryResult> { files::table .filter(files::user_id.eq(user_id)) - .filter(files::session_id.eq(session_id)) + .filter(files::session_id.is_null()) .filter(files::id.eq(file_id)) .select(ChatRsFile::as_select()) .first(self.db) .await + .optional() } - pub async fn list_session_files( + pub async fn find_session_file( &mut self, user_id: &Uuid, session_id: &Uuid, - ) -> QueryResult> { + file_id: &Uuid, + ) -> QueryResult> { files::table .filter(files::user_id.eq(user_id)) .filter(files::session_id.eq(session_id)) + .filter(files::id.eq(file_id)) .select(ChatRsFile::as_select()) - .load(self.db) + .first(self.db) .await + .optional() } pub async fn attach_files( @@ -73,6 +73,53 @@ impl<'a> FileRepository<'a> { .await } + pub async fn list_user_files(&mut self, user_id: &Uuid) -> QueryResult> { + files::table + .filter(files::user_id.eq(user_id)) + .filter(files::session_id.is_null()) + .select(ChatRsFile::as_select()) + .load(self.db) + .await + } + + pub async fn list_session_files_and_attachments( + &mut self, + user_id: &Uuid, + session_id: &Uuid, + ) -> QueryResult<(Vec, Vec)> { + let (files, attachments) = futures::future::try_join( + files::table + .filter(files::user_id.eq(user_id)) + .filter(files::session_id.eq(session_id)) + .select(ChatRsFile::as_select()) + .load(&mut self.db.as_ref()), + files::table + .inner_join(message_attachments::table) + .inner_join( + chat_messages::table.on(message_attachments::message_id.eq(chat_messages::id)), + ) + .inner_join( + chat_sessions::table.on(chat_sessions::id.eq(chat_messages::session_id)), + ) + .filter(chat_sessions::user_id.eq(user_id)) + .filter(chat_sessions::id.eq(session_id)) + .select(ChatRsMessageAttachment::as_select()) + .load(&mut self.db.as_ref()), + ) + .await?; + + Ok((files, attachments)) + } + + pub async fn delete_user_file(&mut self, user_id: &Uuid, file_id: &Uuid) -> QueryResult { + diesel::delete(files::table) + .filter(files::user_id.eq(user_id)) + .filter(files::id.eq(file_id)) + .returning(files::id) + .get_result(self.db) + .await + } + pub async fn delete_session_file( &mut self, user_id: &Uuid, diff --git a/server-new/src/services/storage/engines.rs b/server-new/src/services/storage/engines.rs index 4b15657..e9d58fc 100644 --- a/server-new/src/services/storage/engines.rs +++ b/server-new/src/services/storage/engines.rs @@ -1,3 +1,5 @@ +//! Storage engines + mod local; pub use local::LocalStorage; diff --git a/server-new/src/services/storage/engines/local.rs b/server-new/src/services/storage/engines/local.rs index 805b214..9ce864f 100644 --- a/server-new/src/services/storage/engines/local.rs +++ b/server-new/src/services/storage/engines/local.rs @@ -1,6 +1,6 @@ use std::path::{Path, PathBuf}; -use futures::future::BoxFuture; +use futures::{FutureExt, future::BoxFuture}; use tokio::{ fs::File, io::{AsyncRead, AsyncReadExt, AsyncWriteExt, BufWriter}, @@ -17,14 +17,12 @@ pub struct LocalStorage { } impl LocalStorage { - pub fn new(base_path: &str) -> Self { - Self { - base_path: PathBuf::from(base_path), - } + pub fn new(base_path: PathBuf) -> Self { + Self { base_path } } - fn file_path(&self, path: &Path) -> PathBuf { - self.base_path.join(path) + fn local_path(&self, file_path: &Path) -> PathBuf { + self.base_path.join(file_path) } async fn file_exists(path: &Path) -> bool { @@ -37,55 +35,46 @@ impl LocalStorage { impl StorageEngine for LocalStorage { fn create<'r>( - &self, - path: &Path, + &'r self, + file_path: &'r Path, reader: &'r mut (dyn AsyncRead + Unpin + Send), ) -> BoxFuture<'r, StorageResult> { - let file_path = self.file_path(path); - Box::pin(async move { - let dir = file_path.parent().expect("should have a parent directory"); + let local_path = self.local_path(file_path); + let dir = local_path.parent().expect("should always have parent dir"); tokio::fs::create_dir_all(&dir).await?; - let mut file = File::create_new(&file_path).await?; + let mut file = File::create_new(&local_path).await?; let mut file_writer = BufWriter::new(&mut file); - let mut read_buffer = [0; 4096]; - let mut total_bytes_written: usize = 0; + let mut read_buffer = [0; 8192]; + let mut total_bytes: usize = 0; - loop { - let n = reader.read(&mut read_buffer).await?; - if n == 0 { - break; - } + while let n = reader.read(&mut read_buffer).await? + && n != 0 + { file_writer.write_all(&read_buffer[..n]).await?; - total_bytes_written += n; + total_bytes += n; } file_writer.flush().await?; file.sync_all().await?; - Ok(total_bytes_written) + Ok(total_bytes) }) } - fn exists(&self, path: &Path) -> BoxFuture<'_, StorageResult> { - let file_path = self.file_path(path); - - Box::pin(async move { Ok(Self::file_exists(&file_path).await) }) + fn exists<'r>(&'r self, file_path: &'r Path) -> BoxFuture<'r, StorageResult> { + async move { Ok(Self::file_exists(&self.local_path(file_path)).await) }.boxed() } - fn delete(&self, path: &Path) -> BoxFuture<'_, StorageResult<()>> { - let file_path = self.file_path(path); - + fn delete<'r>(&'r self, file_path: &'r Path) -> BoxFuture<'r, StorageResult<()>> { Box::pin(async move { - match Self::file_exists(&file_path).await { - true => Ok(tokio::fs::remove_file(&file_path).await?), + let local_path = self.local_path(file_path); + + match Self::file_exists(&local_path).await { + true => Ok(tokio::fs::remove_file(&local_path).await?), false => Err(StorageError::NotFound), } }) } - - fn signed_url(&self, path: &Path) -> StorageResult { - todo!() - } } diff --git a/server-new/src/services/storage/error.rs b/server-new/src/services/storage/error.rs index 1afcb91..bba539a 100644 --- a/server-new/src/services/storage/error.rs +++ b/server-new/src/services/storage/error.rs @@ -1,9 +1,39 @@ +//! Storage operation errors + +use crate::{db::DbPoolError, error::AppError}; + pub type StorageResult = Result; #[derive(Debug, thiserror::Error)] pub enum StorageError { - #[error("IO error: {0}")] - Io(#[from] std::io::Error), #[error("File not found")] NotFound, + #[error("File with this path already exists")] + AlreadyExists, + #[error("File is missing content type")] + MissingContentType, + #[error("Unsupported content type: '{0}'")] + UnsupportedContentType(String), + #[error("Invalid file name/path: '{0}'")] + InvalidPath(String), + + #[error("IO error: {0}")] + Io(#[from] std::io::Error), + #[error("Database error: {0}")] + Database(#[from] diesel::result::Error), + #[error(transparent)] + DatabasePool(#[from] DbPoolError), +} + +impl From for AppError { + fn from(error: StorageError) -> Self { + match error { + StorageError::NotFound => AppError::not_found(error.to_string()), + StorageError::AlreadyExists + | StorageError::MissingContentType + | StorageError::UnsupportedContentType(_) + | StorageError::InvalidPath(_) => Self::bad_request(error.to_string()), + err => Self::internal(err.into()), + } + } } diff --git a/server-new/src/services/storage/interface.rs b/server-new/src/services/storage/interface.rs index 96ebe69..007c3dd 100644 --- a/server-new/src/services/storage/interface.rs +++ b/server-new/src/services/storage/interface.rs @@ -1,3 +1,5 @@ +//! Storage interface + use std::path::Path; use futures::future::BoxFuture; @@ -8,11 +10,13 @@ use super::error::StorageResult; /// Trait representing an underlying storage to manage files for LLM chats and responses pub trait StorageEngine: Send + Sync { fn create<'r>( - &self, - path: &Path, + &'r self, + path: &'r Path, reader: &'r mut (dyn AsyncRead + Unpin + Send), ) -> BoxFuture<'r, StorageResult>; - fn exists(&self, path: &Path) -> BoxFuture<'_, StorageResult>; - fn delete(&self, path: &Path) -> BoxFuture<'_, StorageResult<()>>; - fn signed_url(&self, path: &Path) -> StorageResult; + fn exists<'r>(&'r self, path: &'r Path) -> BoxFuture<'r, StorageResult>; + fn delete<'r>(&'r self, path: &'r Path) -> BoxFuture<'r, StorageResult<()>>; + fn signed_url(&self, #[allow(unused)] path: &Path) -> StorageResult> { + Ok(None) + } } diff --git a/server-new/src/services/storage/mod.rs b/server-new/src/services/storage/mod.rs index 0c2c0b3..59d471e 100644 --- a/server-new/src/services/storage/mod.rs +++ b/server-new/src/services/storage/mod.rs @@ -1,11 +1,19 @@ -use std::path::PathBuf; +use std::path::{Path, PathBuf}; -use axum::extract::multipart::MultipartError; -use futures::{Stream, TryStreamExt}; +use futures::Stream; use tokio_util::io::StreamReader; use uuid::Uuid; -use crate::services::storage::{engines::LocalStorage, error::StorageResult}; +use crate::{ + db::{ + DbPool, DbService, + models::{ChatRsFile, ChatRsFileType, NewChatRsFile}, + }, + services::storage::{ + engines::LocalStorage, + error::{StorageError, StorageResult}, + }, +}; mod engines; mod error; @@ -13,38 +21,136 @@ mod interface; pub use interface::StorageEngine; +/// Name of the base folder containing all files/attachments +pub const STORAGE_FOLDER: &str = "rs-chat/storage"; +/// Name of the folder containing user files +pub const USER_FOLDER: &str = "user"; +/// Name of the folder containing session files +pub const SESSION_FOLDER: &str = "session"; + pub struct StorageService<'r> { - data_dir: &'r str, + data_dir: &'r Path, + db_pool: &'r DbPool, } impl<'r> StorageService<'r> { - fn file_path(&self, user_id: &Uuid, session_id: &Uuid, name: &str) -> PathBuf { - let session_folder = PathBuf::from(format!("{user_id}/{session_id}")); - - session_folder.join(name) + pub fn new(data_dir: &'r Path, db_pool: &'r DbPool) -> Self { + Self { data_dir, db_pool } } pub async fn create_file( &self, user_id: &Uuid, - session_id: &Uuid, - name: &str, - stream: impl Stream> + Send + Unpin, - ) -> StorageResult<()> { - let storage: Box = - Box::new(LocalStorage::new(&format!("{}/storage", self.data_dir))); - - let path = self.file_path(user_id, session_id, name); - let mut reader = StreamReader::new(stream.map_err(std::io::Error::other)); - - let size = match storage.create(&path, &mut reader).await { - Ok(n) => n, + session_id: Option<&Uuid>, + path: &str, + content_type: Option<&str>, + stream: impl Stream> + Send + Unpin, + ) -> StorageResult { + let file_path = self.build_file_path(user_id, session_id, &path)?; + let content_type = content_type.ok_or(StorageError::MissingContentType)?; + let file_type = self.validate_file_type(content_type)?; + + let storage = self.storage_engine(); + if storage.exists(&file_path).await? { + return Err(StorageError::AlreadyExists); + } + + let mut reader = StreamReader::new(stream); + let file_size = match storage.create(&file_path, &mut reader).await { + Ok(num_bytes) => num_bytes, Err(err) => { - let _ = storage.delete(&path).await; + let _ = storage.delete(&file_path).await; return Err(err); } }; - todo!() + let mut db = DbService::from_pool(self.db_pool).await?; + let new_file = NewChatRsFile { + user_id, + session_id, + path, + file_type: file_type.as_ref(), + content_type, + size: file_size.try_into().unwrap_or_default(), + }; + let db_file = db.files().create_file(new_file).await?; + + Ok(db_file) + } + + pub async fn delete_file( + &self, + db: &mut DbService, + user_id: &Uuid, + session_id: Option<&Uuid>, + file_id: &Uuid, + ) -> StorageResult { + let db_file = match session_id { + Some(session_id) => { + db.files() + .find_session_file(user_id, session_id, file_id) + .await? + } + None => db.files().find_user_file(user_id, file_id).await?, + } + .ok_or(StorageError::NotFound)?; + + let storage = self.storage_engine(); + let file_path = self.build_file_path(user_id, session_id, &db_file.path)?; + if let Err(err) = storage.delete(&file_path).await { + tracing::warn!("error deleting file {file_id} with path {file_path:?}: {err}"); + } + + let deleted_file_id = match session_id { + Some(session_id) => { + db.files() + .delete_session_file(user_id, session_id, file_id) + .await? + } + None => db.files().delete_user_file(user_id, file_id).await?, + }; + + Ok(deleted_file_id) + } + + fn storage_engine(&self) -> Box { + Box::new(LocalStorage::new(self.data_dir.join(STORAGE_FOLDER))) + } + + fn build_file_path( + &self, + user_id: &Uuid, + session_id: Option<&Uuid>, + path: &str, + ) -> StorageResult { + let valid_path = Path::new(path).is_relative() + && Path::new(path) + .components() + .all(|c| matches!(c, std::path::Component::Normal(_))); + if !valid_path { + return Err(StorageError::InvalidPath(path.into())); + } + + Ok(match session_id { + Some(session_id) => { + let segments = [ + &user_id.to_string(), + SESSION_FOLDER, + &session_id.to_string(), + &path, + ]; + segments.iter().collect() + } + None => [&user_id.to_string(), USER_FOLDER, &path].iter().collect(), + }) + } + + fn validate_file_type(&self, content_type: &str) -> StorageResult { + match content_type { + "image/jpeg" | "image/png" | "image/webp" => Ok(ChatRsFileType::Image), + "application/pdf" => Ok(ChatRsFileType::Pdf), + text if text.starts_with("text/") => Ok(ChatRsFileType::Text), + unsupported => Err(StorageError::UnsupportedContentType(unsupported.into())), + } } } diff --git a/server-new/src/state.rs b/server-new/src/state.rs index c58708a..a8baf1a 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -12,6 +12,7 @@ use crate::{ chat::ChatService, model::ModelService, provider::ProviderService, + storage::StorageService, stream::tinistream::TinistreamClient, }, }; @@ -44,6 +45,9 @@ impl AppState { pub fn model_service(&self) -> ModelService<'_> { ModelService::new(&self.redis, &self.http_client) } + pub fn storage_service(&self) -> StorageService<'_> { + StorageService::new(&self.config.server.data_dir, &self.db_pool) + } } impl Deref for AppState { From 4f35bf3f23cfa441af657b0c1347e1d51fed5778 Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sat, 18 Jul 2026 16:11:44 -0400 Subject: [PATCH 108/111] local file storage: clean up folders when deleting files --- .../src/services/storage/engines/local.rs | 29 ++++++++++++------- 1 file changed, 18 insertions(+), 11 deletions(-) diff --git a/server-new/src/services/storage/engines/local.rs b/server-new/src/services/storage/engines/local.rs index 9ce864f..59ee490 100644 --- a/server-new/src/services/storage/engines/local.rs +++ b/server-new/src/services/storage/engines/local.rs @@ -24,13 +24,6 @@ impl LocalStorage { fn local_path(&self, file_path: &Path) -> PathBuf { self.base_path.join(file_path) } - - async fn file_exists(path: &Path) -> bool { - match tokio::fs::metadata(&path).await { - Ok(meta) => meta.is_file(), - Err(_) => false, - } - } } impl StorageEngine for LocalStorage { @@ -64,17 +57,31 @@ impl StorageEngine for LocalStorage { } fn exists<'r>(&'r self, file_path: &'r Path) -> BoxFuture<'r, StorageResult> { - async move { Ok(Self::file_exists(&self.local_path(file_path)).await) }.boxed() + async move { Ok(tokio::fs::try_exists(&self.local_path(file_path)).await?) }.boxed() } fn delete<'r>(&'r self, file_path: &'r Path) -> BoxFuture<'r, StorageResult<()>> { Box::pin(async move { let local_path = self.local_path(file_path); + if tokio::fs::try_exists(&local_path).await? { + return Err(StorageError::NotFound); + } - match Self::file_exists(&local_path).await { - true => Ok(tokio::fs::remove_file(&local_path).await?), - false => Err(StorageError::NotFound), + tokio::fs::remove_file(&local_path).await?; + + // Clean up parent directories + let mut parent_dir = local_path.clone(); + while let Some(dir) = parent_dir.parent().filter(|dir| *dir != self.base_path) { + match tokio::fs::remove_dir(dir).await { + Ok(_) => parent_dir = dir.to_path_buf(), + Err(err) if err.kind() == std::io::ErrorKind::DirectoryNotEmpty => { + break; + } + Err(err) => return Err(StorageError::Io(err)), + }; } + + Ok(()) }) } } From 7cead1e4a715051c9ae61c5fdb9d8adce801a8cc Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sat, 18 Jul 2026 16:18:02 -0400 Subject: [PATCH 109/111] Update local.rs --- server-new/src/services/storage/engines/local.rs | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/server-new/src/services/storage/engines/local.rs b/server-new/src/services/storage/engines/local.rs index 59ee490..7a4edbe 100644 --- a/server-new/src/services/storage/engines/local.rs +++ b/server-new/src/services/storage/engines/local.rs @@ -1,6 +1,6 @@ use std::path::{Path, PathBuf}; -use futures::{FutureExt, future::BoxFuture}; +use futures::future::BoxFuture; use tokio::{ fs::File, io::{AsyncRead, AsyncReadExt, AsyncWriteExt, BufWriter}, @@ -57,7 +57,11 @@ impl StorageEngine for LocalStorage { } fn exists<'r>(&'r self, file_path: &'r Path) -> BoxFuture<'r, StorageResult> { - async move { Ok(tokio::fs::try_exists(&self.local_path(file_path)).await?) }.boxed() + Box::pin(async move { + let exists = tokio::fs::try_exists(&self.local_path(file_path)).await?; + + Ok(exists) + }) } fn delete<'r>(&'r self, file_path: &'r Path) -> BoxFuture<'r, StorageResult<()>> { From 554f10cc5597078edf9cce83af84541257a680fc Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sat, 18 Jul 2026 17:29:48 -0400 Subject: [PATCH 110/111] update deps --- server-new/Cargo.lock | 187 ++++++++++++++++++++++-------------------- server-new/Cargo.toml | 10 +-- 2 files changed, 104 insertions(+), 93 deletions(-) diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index f5a486d..7cc435d 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -110,7 +110,7 @@ dependencies = [ "darling 0.23.0", "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -124,9 +124,9 @@ dependencies = [ [[package]] name = "anyhow" -version = "1.0.103" +version = "1.0.104" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2a4385e2e34eb35d6b3efe798b9eb88096925d87726c0798709bf56d9ed84af3" +checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470" [[package]] name = "arc-swap" @@ -156,18 +156,18 @@ checksum = "c7c24de15d275a1ecfd47a380fb4d5ec9bfe0933f309ed5e705b775596a3574d" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] name = "async-trait" -version = "0.1.89" +version = "0.1.90" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9035ad2d096bed7955a320ee7e2230574d28fd3c3a0f186cbea1ff3c7eed5dbb" +checksum = "62a5e99d6b2764d521fa86b22ca32ad96f19ae2427febfd80a131a2e3e9d6ad9" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.0", ] [[package]] @@ -209,9 +209,9 @@ checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" [[package]] name = "aws-lc-rs" -version = "1.17.1" +version = "1.17.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4342d8937fc7e5dd9b1c60292261c0670c882a2cd1719cfc11b1af41731e32ad" +checksum = "00bdb5da18dac48ca2cc7cd4a98e533e8635a58e2361d13a1a4ee3888e0d72f1" dependencies = [ "aws-lc-sys", "zeroize", @@ -219,9 +219,9 @@ dependencies = [ [[package]] name = "aws-lc-sys" -version = "0.42.0" +version = "0.43.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6d9ceb1da931507a12f4fccea479dccd00da1943e1b4ae72d8e502d707361444" +checksum = "43103168cc76fe62678a375e722fc9cb3a0146159ac5828bc4f0dfd755c2224c" dependencies = [ "cc", "cmake", @@ -271,7 +271,7 @@ source = "git+https://git.fasharp.io/fa-sharp/axum-aide-macros?rev=5b00e645df#5b dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -327,7 +327,7 @@ source = "git+https://git.fasharp.io/fa-sharp/axum-plugin?rev=be17dc9aec#be17dc9 dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -353,9 +353,9 @@ dependencies = [ [[package]] name = "bitflags" -version = "2.13.0" +version = "2.13.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b4388bee8683e3d04af747c73422af53102d2bd24d9eadb6cbc100baef4b43f8" +checksum = "b588b76d00fde79687d7646a9b5bdf3cc0f655e0bbd080335a95d7e96f3587da" [[package]] name = "block-buffer" @@ -397,7 +397,7 @@ dependencies = [ "proc-macro2", "quote", "rustversion", - "syn", + "syn 2.0.119", ] [[package]] @@ -436,9 +436,9 @@ dependencies = [ [[package]] name = "cc" -version = "1.2.67" +version = "1.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e17dd265a7d0f31ef544e1b20e03add05d3b45b491b633b10d67145d2acc1a38" +checksum = "c89588d05638b5b4594a3348a2d6c20277e43a7f5c5202b05cc56888475a47b8" dependencies = [ "find-msvc-tools", "jobserver", @@ -454,9 +454,9 @@ checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" [[package]] name = "cfg_aliases" -version = "0.2.1" +version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" +checksum = "f079e83a288787bcd14a6aea84cee5c87a67c5a3e660c30f557a3d24761b3527" [[package]] name = "chacha20" @@ -700,7 +700,7 @@ dependencies = [ "proc-macro2", "quote", "strsim", - "syn", + "syn 2.0.119", ] [[package]] @@ -713,7 +713,7 @@ dependencies = [ "proc-macro2", "quote", "strsim", - "syn", + "syn 2.0.119", ] [[package]] @@ -724,7 +724,7 @@ checksum = "d38308df82d1080de0afee5d069fa14b0326a88c14f15c5ccda35b4a6c414c81" dependencies = [ "darling_core 0.21.3", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -735,7 +735,7 @@ checksum = "ac3984ec7bd6cfa798e62b4a642426a5be0e68f9401cfc2a01e3fa9ea2fcdb8d" dependencies = [ "darling_core 0.23.0", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -815,7 +815,7 @@ dependencies = [ "heck 0.4.1", "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -824,7 +824,7 @@ version = "0.1.0" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -837,7 +837,7 @@ dependencies = [ "dsl_auto_type", "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -857,7 +857,7 @@ version = "0.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fe2444076b48641147115697648dc743c2c00b61adade0f01ce67133c7babe8c" dependencies = [ - "syn", + "syn 2.0.119", ] [[package]] @@ -891,7 +891,7 @@ checksum = "1ac70aa55017e108007fbaf5aa0f54b021c98f92ff8af59d42eda9da96e3dd4f" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -917,7 +917,7 @@ dependencies = [ "heck 0.5.0", "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -1047,7 +1047,7 @@ checksum = "1458c6e22d36d61507034d5afecc64f105c1d39712b7ac6ec3b352c423f715cc" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -1058,9 +1058,9 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" [[package]] name = "futures" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8b147ee9d1f6d097cef9ce628cd2ee62288d963e16fb287bd9286455b241382d" +checksum = "a88cf1f829d945f548cf8fec32c61b1f202b6d93b45848602fc02af4b12ad218" dependencies = [ "futures-channel", "futures-core", @@ -1073,9 +1073,9 @@ dependencies = [ [[package]] name = "futures-channel" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "07bbe89c50d7a535e539b8c17bc0b49bdb77747034daa8087407d655f3f7cc1d" +checksum = "262590f4fe6afeb0bc83be1daa64e52657fe185690a958af7f3ad0e92085c5ae" dependencies = [ "futures-core", "futures-sink", @@ -1083,15 +1083,15 @@ dependencies = [ [[package]] name = "futures-core" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d" +checksum = "2cd50c473c80f6d7c3670a752354b8e569b1a7cbfdc0419ec88e5edad85e0dc7" [[package]] name = "futures-executor" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "baf29c38818342a3b26b5b923639e7b1f4a61fc5e76102d4b1981c6dc7a7579d" +checksum = "6754879cc9f2c66f88c6e5c35344bb0bdb0708b0352b1201815667c7eabc7458" dependencies = [ "futures-core", "futures-task", @@ -1100,38 +1100,38 @@ dependencies = [ [[package]] name = "futures-io" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cecba35d7ad927e23624b22ad55235f2239cfa44fd10428eecbeba6d6a717718" +checksum = "4577ecaa3c4f96589d473f679a71b596316f6641bc350038b962a5daf0085d7a" [[package]] name = "futures-macro" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e835b70203e41293343137df5c0664546da5745f82ec9b84d40be8336958447b" +checksum = "2d6d3cde68c518367be28956066ddfef33813991b77a55005a69dae04bf3b10b" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] name = "futures-sink" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c39754e157331b013978ec91992bde1ac089843443c49cbc7f46150b0fad0893" +checksum = "e34418ac499d6305c2fb5ad0ed2f6ac998c5f8ca209b4510f7f94242c647e307" [[package]] name = "futures-task" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "037711b3d59c33004d3856fbdc83b99d4ff37a24768fa1be9ce3538a1cde4393" +checksum = "b231ed28831efb4a61a08580c4bc233ec56bc009f4cd8f52da2c3cb97df0c109" [[package]] name = "futures-util" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" +checksum = "a77a90a256fce34da66415271e30f94ee91c57b04b8a2c042d9cf3220179deaa" dependencies = [ "futures-channel", "futures-core", @@ -1209,7 +1209,7 @@ version = "0.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2eecf2d5dc9b66b732b97707a0210906b1d30523eb773193ab777c0c84b3e8d5" dependencies = [ - "polyval 0.7.2", + "polyval 0.7.3", ] [[package]] @@ -1602,7 +1602,7 @@ dependencies = [ "quote", "rustc_version", "simd_cesu8", - "syn", + "syn 2.0.119", ] [[package]] @@ -1621,7 +1621,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264" dependencies = [ "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -1972,7 +1972,7 @@ dependencies = [ "proc-macro2", "proc-macro2-diagnostics", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -2026,9 +2026,9 @@ dependencies = [ [[package]] name = "polyval" -version = "0.7.2" +version = "0.7.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b20f20e954175de5f463f67781b35583397d916b1d148738923711b2ad16bee8" +checksum = "f0fa31d631f2b2cb2a544d0aa321ce847a94764d701ca2becc411138b93d49cd" dependencies = [ "cpubits", "cpufeatures 0.3.0", @@ -2095,7 +2095,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" dependencies = [ "proc-macro2", - "syn", + "syn 2.0.119", ] [[package]] @@ -2115,7 +2115,7 @@ checksum = "af066a9c399a26e020ada66a034357a868728e72cd426f3adcd35f80d88d88c8" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", "version_check", "yansi", ] @@ -2338,14 +2338,14 @@ checksum = "b7186006dcb21920990093f30e3dea63b7d6e977bf1256be20c3563a5db070da" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] name = "regex-automata" -version = "0.4.15" +version = "0.4.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1f388202e4b80542a0921078cc23b6333bcf1409c1e3f86404cae4766a6131db" +checksum = "8fcfdb36bda0c880c5931cdc7a2bcdc8ba4556847b9d912bca70bc94708711ad" dependencies = [ "aho-corasick", "memchr", @@ -2645,7 +2645,7 @@ dependencies = [ "proc-macro2", "quote", "serde_derive_internals", - "syn", + "syn 2.0.119", ] [[package]] @@ -2710,7 +2710,7 @@ checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -2721,7 +2721,7 @@ checksum = "18d26a20a969b9e3fdf2fc2d9f21eda6c40e2de84c9408bb5d3b05d499aae711" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -2811,7 +2811,7 @@ dependencies = [ "darling 0.23.0", "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -2988,7 +2988,7 @@ dependencies = [ "heck 0.5.0", "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -3005,9 +3005,20 @@ checksum = "a7973cce6668464ea31f176d85b13c7ab3bba2cb3b77a2ed26abd7801688010a" [[package]] name = "syn" -version = "2.0.118" +version = "2.0.119" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "872831b642d1a07999a962a351ed35b955ea2cfc8f3862091e2a240a84f17297" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "syn" +version = "3.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1b9ae57f904213ebb649ce6895b8a66c66f0203b9319718f69a5612a065b1422" +checksum = "f2fac314a64dc9a36e61a9eb4261a5e9bbfbc922b27e518af97bc32b926cf967" dependencies = [ "proc-macro2", "quote", @@ -3031,7 +3042,7 @@ checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -3060,7 +3071,7 @@ checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -3071,7 +3082,7 @@ checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -3153,9 +3164,9 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokio" -version = "1.52.3" +version = "1.53.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe" +checksum = "d988bcd52dbe076d3d46903332f58c912b87a2c49b1428419a5845154762ffee" dependencies = [ "bytes", "libc", @@ -3169,13 +3180,13 @@ dependencies = [ [[package]] name = "tokio-macros" -version = "2.7.0" +version = "2.7.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496" +checksum = "6328af13490e73a9b4694030fafd93f8c8c6a9dede33e821c3fc63eddf8042ba" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -3497,7 +3508,7 @@ checksum = "7490cfa5ec963746568740651ac6781f701c9c5ea257c58e057f3ba8cf69e8da" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -3691,9 +3702,9 @@ checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" [[package]] name = "uuid" -version = "1.23.5" +version = "1.24.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ea5fab0d6c3c01ae70085a09cb03d4c7a1d6314e2b3e075392783396d724ca0a" +checksum = "bf3923a6f5c4c6382e0b653c4117f48d631ea17f38ed86e2a828e6f7412f5239" dependencies = [ "getrandom 0.4.3", "js-sys", @@ -3807,7 +3818,7 @@ dependencies = [ "bumpalo", "proc-macro2", "quote", - "syn", + "syn 2.0.119", "wasm-bindgen-shared", ] @@ -3855,9 +3866,9 @@ dependencies = [ [[package]] name = "webpki-root-certs" -version = "1.0.8" +version = "1.0.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0d46a5a140e6f7afeccd8eae97eff335163939eac8b929834875168b29b3d267" +checksum = "b96554aa2acc8ccdb7e1c9a58a7a68dd5d13bccc69cd124cb09406db612a1c9b" dependencies = [ "rustls-pki-types", ] @@ -3905,7 +3916,7 @@ checksum = "053e2e040ab57b9dc951b72c264860db7eb3b0200ba345b4e4c3b14f67855ddf" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -3916,7 +3927,7 @@ checksum = "3f316c4a2570ba26bbec722032c4099d8c8bc095efccdc15688708623367e358" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -4077,7 +4088,7 @@ checksum = "de844c262c8848816172cef550288e7dc6c7b7814b4ee56b3e1553f275f1858e" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", "synstructure", ] @@ -4098,7 +4109,7 @@ checksum = "e2e817b7b52d0c7358d3246da9d69935ebb18116b2b102b4230dac079b4862f5" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -4118,7 +4129,7 @@ checksum = "11532158c46691caf0f2593ea8358fed6bbf68a0315e80aae9bd41fbade684a1" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", "synstructure", ] @@ -4158,7 +4169,7 @@ checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index d5ae512..fa504c7 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -19,9 +19,9 @@ aide = { "swagger" ] } -anyhow = "1.0.103" +anyhow = "1.0.104" async-stream = "0.3.6" -async-trait = "0.1.89" +async-trait = "0.1.90" axum = { version = "0.8.9", features = ["json", "multipart", "query"] } axum-aide-macros = { git = "https://git.fasharp.io/fa-sharp/axum-aide-macros", @@ -57,7 +57,7 @@ fred = { default-features = false, features = ["i-keys", "i-streams"] } -futures = "0.3.32" +futures = "0.3.33" hex = "0.4.3" reqwest = { version = "0.13.4", @@ -88,7 +88,7 @@ tinistream-client = { rev = "015d307" } tokio = { - version = "1.52.3", + version = "1.53.0", default-features = false, features = ["macros", "net", "rt", "rt-multi-thread", "signal"] } @@ -111,4 +111,4 @@ tower-sessions-redis-store = { tracing = "0.1.44" tracing-appender = "0.2.5" tracing-subscriber = { version = "0.3.23", features = ["env-filter", "json"] } -uuid = { version = "1.23.5", features = ["serde", "v4"] } +uuid = { version = "1.24.0", features = ["serde", "v4"] } From 8ef9262b18b9779277f35e1d7a3850690e7d319a Mon Sep 17 00:00:00 2001 From: fa-sharp Date: Sun, 19 Jul 2026 00:59:19 -0400 Subject: [PATCH 111/111] s3 upload --- server-new/Cargo.lock | 195 ++++++++++++++---- server-new/Cargo.toml | 15 +- server-new/src/api/upload.rs | 43 ++-- server-new/src/config.rs | 20 +- server-new/src/extractors/mod.rs | 2 + server-new/src/extractors/upload.rs | 57 +++++ server-new/src/plugins/security.rs | 6 +- server-new/src/services/storage/engines.rs | 2 + .../src/services/storage/engines/local.rs | 10 +- server-new/src/services/storage/engines/s3.rs | 154 ++++++++++++++ server-new/src/services/storage/error.rs | 11 +- server-new/src/services/storage/interface.rs | 7 +- server-new/src/services/storage/mod.rs | 58 ++++-- server-new/src/state.rs | 7 +- 14 files changed, 482 insertions(+), 105 deletions(-) create mode 100644 server-new/src/extractors/upload.rs create mode 100644 server-new/src/services/storage/engines/s3.rs diff --git a/server-new/Cargo.lock b/server-new/Cargo.lock index 7cc435d..29f06a6 100644 --- a/server-new/Cargo.lock +++ b/server-new/Cargo.lock @@ -249,7 +249,6 @@ dependencies = [ "matchit", "memchr", "mime", - "multer", "percent-encoding", "pin-project-lite", "serde_core", @@ -293,6 +292,29 @@ dependencies = [ "tracing", ] +[[package]] +name = "axum-extra" +version = "0.12.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "be44683b41ccb9ab2d23a5230015c9c3c55be97a25e4428366de8873103f7970" +dependencies = [ + "axum", + "axum-core", + "bytes", + "futures-core", + "futures-util", + "http", + "http-body", + "http-body-util", + "mime", + "pin-project-lite", + "tokio", + "tokio-util", + "tower-layer", + "tower-service", + "tracing", +] + [[package]] name = "axum-helmet" version = "1.0.2" @@ -351,6 +373,12 @@ dependencies = [ "serde_json", ] +[[package]] +name = "bitflags" +version = "1.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bef38d45163c2f1dde094a7dfd33ccf595c92905c8f8f4fdc18d06fb1037718a" + [[package]] name = "bitflags" version = "2.13.1" @@ -761,6 +789,37 @@ version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2657f61fb1dd8bf37a8d51093cc7cee4e77125b22f7753f49b289f831bec2bae" +[[package]] +name = "defmt" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e2953bfe4f93bbd20cc71198842756f77d161884c99ebbabc41d80231ded88d1" +dependencies = [ + "bitflags 1.3.2", + "defmt-macros", +] + +[[package]] +name = "defmt-macros" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bad9c72e7ca2137e0dc3813245a0d282fd6daad32fd800af018306a9169b5fe8" +dependencies = [ + "defmt-parser", + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "defmt-parser" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10d60334b3b2e7c9d91ef8150abfb6fa4c1c39ebbcf4a81c2e346aad939fee3e" +dependencies = [ + "thiserror 2.0.18", +] + [[package]] name = "deranged" version = "0.5.8" @@ -777,7 +836,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e54d1f576cd3a3460f212a4615fd12ce1b6303c095b79a44449ffbe627753dc1" dependencies = [ "bigdecimal", - "bitflags", + "bitflags 2.13.1", "byteorder", "chrono", "diesel_derives", @@ -938,15 +997,6 @@ version = "1.16.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" -[[package]] -name = "encoding_rs" -version = "0.8.35" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "75030f3c4f45dafd7586dd6780965a8c7e8e285a5ecb86713e63a79c5b2766f3" -dependencies = [ - "cfg-if", -] - [[package]] name = "equivalent" version = "1.0.2" @@ -1563,6 +1613,29 @@ dependencies = [ "hybrid-array", ] +[[package]] +name = "instant-xml" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0a2cad967c3b727c000ebfdcd14974539e8c6f59e1d044c8034179df2c6fe250" +dependencies = [ + "instant-xml-macros", + "thiserror 2.0.18", + "xmlparser", +] + +[[package]] +name = "instant-xml-macros" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "44127a3a387c070ef0656a6ce53dd0e616cf8d6cf5b159aa478cfd49e1c166e0" +dependencies = [ + "heck 0.5.0", + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "ipnet" version = "2.12.0" @@ -1575,6 +1648,31 @@ version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" +[[package]] +name = "jiff" +version = "0.2.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "961d16382652bfdd8c6f68b223b26a8c93e0d475c672f414411db31c6c5c900e" +dependencies = [ + "defmt", + "jiff-static", + "log", + "portable-atomic", + "portable-atomic-util", + "serde_core", +] + +[[package]] +name = "jiff-static" +version = "0.2.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d0879bd39df99c4c5e2c6615ccc026391a423dde10532c573e6086eb94a802cc" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "jni" version = "0.22.4" @@ -1784,23 +1882,6 @@ dependencies = [ "windows-sys 0.61.2", ] -[[package]] -name = "multer" -version = "3.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "83e87776546dc87511aa5ee218730c92b666d7264ab6ed41f9d215af9cd5224b" -dependencies = [ - "bytes", - "encoding_rs", - "futures-util", - "http", - "httparse", - "memchr", - "mime", - "spin", - "version_check", -] - [[package]] name = "nom" version = "7.1.3" @@ -1899,7 +1980,7 @@ version = "0.3.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2a180dd8642fa45cdb7dd721cd4c11b1cadd4929ce112ebd8b9f5803cc79d536" dependencies = [ - "bitflags", + "bitflags 2.13.1", ] [[package]] @@ -2035,6 +2116,21 @@ dependencies = [ "universal-hash 0.6.1", ] +[[package]] +name = "portable-atomic" +version = "1.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3d20d5497ef88037a52ff98267d066e7f11fcc5e99bbfbd58a42336193aacec3" + +[[package]] +name = "portable-atomic-util" +version = "0.2.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a106d1259c23fac8e543272398ae0e3c0b8d33c88ed73d0cc71b0f1d902618" +dependencies = [ + "portable-atomic", +] + [[package]] name = "postgres-protocol" version = "0.6.12" @@ -2318,7 +2414,7 @@ version = "0.5.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d" dependencies = [ - "bitflags", + "bitflags 2.13.1", ] [[package]] @@ -2463,6 +2559,7 @@ dependencies = [ "async-trait", "axum", "axum-aide-macros", + "axum-extra", "axum-helmet", "axum-plugin", "bigdecimal", @@ -2478,6 +2575,7 @@ dependencies = [ "hex", "reqwest", "reqwest-websocket", + "rusty-s3", "schemars", "serde", "serde_json", @@ -2595,6 +2693,25 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f" +[[package]] +name = "rusty-s3" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "20f0d23aa8ac3b44d4cfb1e4b3611e6f3776debfb3f7701c4ea9f2252a701403" +dependencies = [ + "base64", + "hmac 0.13.0", + "instant-xml", + "jiff", + "md-5", + "percent-encoding", + "serde", + "serde_json", + "sha2 0.11.0", + "url", + "zeroize", +] + [[package]] name = "ryu" version = "1.0.23" @@ -2660,7 +2777,7 @@ version = "3.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d" dependencies = [ - "bitflags", + "bitflags 2.13.1", "core-foundation", "core-foundation-sys", "libc", @@ -2941,12 +3058,6 @@ dependencies = [ "windows-sys 0.61.2", ] -[[package]] -name = "spin" -version = "0.9.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e" - [[package]] name = "stable_deref_trait" version = "1.2.1" @@ -3360,7 +3471,7 @@ version = "0.6.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840" dependencies = [ - "bitflags", + "bitflags 2.13.1", "bytes", "futures-util", "http", @@ -3378,7 +3489,7 @@ version = "0.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b11f75e912b0c2be01b63d8cf8057b8c3f97cf34abb3d431a3a4c8675498e233" dependencies = [ - "bitflags", + "bitflags 2.13.1", "bytes", "futures-core", "futures-util", @@ -4063,6 +4174,12 @@ version = "0.6.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4" +[[package]] +name = "xmlparser" +version = "0.13.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "66fee0b777b0f5ac1c69bb06d361268faafa61cd4682ae064a171c16c433e9e4" + [[package]] name = "yansi" version = "1.0.1" diff --git a/server-new/Cargo.toml b/server-new/Cargo.toml index fa504c7..eef17ac 100644 --- a/server-new/Cargo.toml +++ b/server-new/Cargo.toml @@ -10,23 +10,17 @@ aes-gcm = "0.11.0" aide = { git = "https://github.com/hniksic/aide.git", rev = "7246c20", - features = [ - "axum", - "axum-json", - "axum-multipart", - "axum-query", - "macros", - "swagger" - ] + features = ["axum", "axum-json", "axum-query", "macros", "swagger"] } anyhow = "1.0.104" async-stream = "0.3.6" async-trait = "0.1.90" -axum = { version = "0.8.9", features = ["json", "multipart", "query"] } +axum = { version = "0.8.9", features = ["json", "query"] } axum-aide-macros = { git = "https://git.fasharp.io/fa-sharp/axum-aide-macros", rev = "5b00e645df" } +axum-extra = { version = "0.12.6", features = ["file-stream"] } axum-helmet = "1.0.2" axum-plugin = { git = "https://git.fasharp.io/fa-sharp/axum-plugin", @@ -65,6 +59,7 @@ reqwest = { features = ["default-tls", "json", "stream"] } reqwest-websocket = { version = "0.6.0", features = ["json"] } +rusty-s3 = "0.10.0" schemars = { version = "1.2.1", features = ["bigdecimal04", "chrono04", "preserve_order", "uuid1"] @@ -97,7 +92,7 @@ tokio-util = { version = "0.7.18", features = ["io"] } tower = { version = "0.5", default-features = false } tower-http = { version = "0.7.0", - features = ["fs", "limit", "request-id", "timeout", "trace"] + features = ["fs", "request-id", "timeout", "trace"] } tower-sessions = { version = "0.15.0", diff --git a/server-new/src/api/upload.rs b/server-new/src/api/upload.rs index 1560afb..86e9c69 100644 --- a/server-new/src/api/upload.rs +++ b/server-new/src/api/upload.rs @@ -1,16 +1,15 @@ use axum::{ Json, - extract::{Multipart, Path, State}, + extract::{Path, State}, }; use axum_aide_macros::api_routes; -use futures::TryStreamExt; use uuid::Uuid; use crate::{ api::ApiTag, db::models::ChatRsFile, error::{AppError, AppResult}, - extractors::{CurrentUser, Database}, + extractors::{CurrentUser, Database, FileUpload}, state::AppState, }; @@ -31,19 +30,18 @@ async fn upload_user_file( CurrentUser { user_id }: CurrentUser, Path(path): Path, State(state): State, - mut multipart: Multipart, + upload: FileUpload, ) -> AppResult> { - let field = multipart - .next_field() - .await - .map_err(|err| AppError::bad_request(err.body_text()))? - .ok_or_else(|| AppError::bad_request("no file in request"))?; - let mime = field.content_type().map(str::to_owned); - let stream = field.map_err(std::io::Error::other); - let file = state .storage_service() - .create_file(&user_id, None, &path, mime.as_deref(), stream) + .create_file( + &user_id, + None, + &path, + upload.size(), + &upload.content_type(), + upload.into_stream(), + ) .await?; Ok(Json(file)) @@ -54,24 +52,23 @@ async fn upload_session_file( Path((sess_id, path)): Path<(Uuid, String)>, Database(mut db): Database, State(state): State, - mut multipart: Multipart, + upload: FileUpload, ) -> AppResult> { if db.chats().find_session(&user_id, &sess_id).await?.is_none() { return Err(AppError::not_found("session not found")); } drop(db); // free database connection since this could be long-running request - let field = multipart - .next_field() - .await - .map_err(|err| AppError::bad_request(err.body_text()))? - .ok_or_else(|| AppError::bad_request("no file in request"))?; - let mime = field.content_type().map(str::to_owned); - let stream = field.map_err(std::io::Error::other); - let file = state .storage_service() - .create_file(&user_id, Some(&sess_id), &path, mime.as_deref(), stream) + .create_file( + &user_id, + Some(&sess_id), + &path, + upload.size(), + &upload.content_type(), + upload.into_stream(), + ) .await?; Ok(Json(file)) diff --git a/server-new/src/config.rs b/server-new/src/config.rs index 5bcbf03..a968025 100644 --- a/server-new/src/config.rs +++ b/server-new/src/config.rs @@ -6,9 +6,12 @@ use axum_plugin::figment::{ }; use serde::{Deserialize, Serialize}; -use crate::services::auth::{ - oauth::{DiscordOAuthConfig, GitHubOAuthConfig, GoogleOAuthConfig, OidcConfig}, - proxy::ProxyHeaderConfig, +use crate::services::{ + auth::{ + oauth::{DiscordOAuthConfig, GitHubOAuthConfig, GoogleOAuthConfig, OidcConfig}, + proxy::ProxyHeaderConfig, + }, + storage::engines::S3Config, }; /// Extract configuration from defaults, local `config.toml`, then `RS_CHAT_` environment variables split by `__`. @@ -28,6 +31,7 @@ pub struct AppConfig { pub services: ServiceConfig, pub security: SecurityConfig, pub redis: RedisConfig, + pub storage: StorageConfig, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -117,7 +121,7 @@ pub struct SecurityConfig { impl Default for SecurityConfig { fn default() -> Self { Self { - body_limit: 5 * 1024 * 1024, // 5 MB + body_limit: 1 * 1024 * 1024, // 1 MB upload_limit: 5 * 1024 * 1024, // 5 MB request_timeout: 120, // 2 minutes } @@ -139,3 +143,11 @@ impl Default for RedisConfig { } } } + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +#[serde(tag = "engine", rename_all = "lowercase")] +pub enum StorageConfig { + #[default] + Local, + S3(S3Config), +} diff --git a/server-new/src/extractors/mod.rs b/server-new/src/extractors/mod.rs index 2871247..7126bd3 100644 --- a/server-new/src/extractors/mod.rs +++ b/server-new/src/extractors/mod.rs @@ -3,9 +3,11 @@ mod auth_config; mod database; mod session; +mod upload; mod user; pub use auth_config::PublicAuthConfig; pub use database::Database; pub use session::{AppSession, SessionMeta}; +pub use upload::FileUpload; pub use user::CurrentUser; diff --git a/server-new/src/extractors/upload.rs b/server-new/src/extractors/upload.rs new file mode 100644 index 0000000..8e89396 --- /dev/null +++ b/server-new/src/extractors/upload.rs @@ -0,0 +1,57 @@ +use aide::OperationIo; +use axum::{RequestExt, extract::FromRequest, http::header}; +use futures::{Stream, TryStreamExt}; + +use crate::{error::AppError, state::AppState}; + +/// Extractor to get a streaming uploaded file +#[derive(OperationIo)] +pub struct FileUpload { + body: axum::body::Body, + content_type: String, + content_length: usize, +} + +impl FromRequest for FileUpload { + type Rejection = AppError; + + async fn from_request( + req: axum::extract::Request, + _state: &AppState, + ) -> Result { + let content_type = req + .headers() + .get(header::CONTENT_TYPE) + .and_then(|h| h.to_str().ok()) + .ok_or_else(|| AppError::bad_request("no content-type header"))? + .to_owned(); + let content_length: usize = req + .headers() + .get(header::CONTENT_LENGTH) + .and_then(|h| h.to_str().ok()) + .and_then(|h| h.parse().ok()) + .ok_or_else(|| AppError::bad_request("no content-length header"))?; + + let body = req.into_limited_body(); + + Ok(Self { + body, + content_length, + content_type, + }) + } +} + +impl FileUpload { + pub fn content_type(&self) -> String { + self.content_type.to_owned() + } + + pub fn size(&self) -> usize { + self.content_length + } + + pub fn into_stream(self) -> impl Stream> { + self.body.into_data_stream().map_err(std::io::Error::other) + } +} diff --git a/server-new/src/plugins/security.rs b/server-new/src/plugins/security.rs index 1e58a77..134f8a5 100644 --- a/server-new/src/plugins/security.rs +++ b/server-new/src/plugins/security.rs @@ -1,8 +1,8 @@ use std::time::Duration; -use axum::http::StatusCode; +use axum::{extract::DefaultBodyLimit, http::StatusCode}; use tower::ServiceBuilder; -use tower_http::{limit::RequestBodyLimitLayer, timeout::TimeoutLayer}; +use tower_http::timeout::TimeoutLayer; use crate::plugins::AxumPlugin; @@ -19,7 +19,7 @@ pub fn plugin() -> AxumPlugin { .into_layer()?; let service = ServiceBuilder::new() - .layer(RequestBodyLimitLayer::new(app.config().security.body_limit)) + .layer(DefaultBodyLimit::max(app.config().security.body_limit)) .layer(TimeoutLayer::with_status_code( StatusCode::REQUEST_TIMEOUT, Duration::from_secs(app.config().security.request_timeout), diff --git a/server-new/src/services/storage/engines.rs b/server-new/src/services/storage/engines.rs index e9d58fc..e15b9f9 100644 --- a/server-new/src/services/storage/engines.rs +++ b/server-new/src/services/storage/engines.rs @@ -1,5 +1,7 @@ //! Storage engines mod local; +mod s3; pub use local::LocalStorage; +pub use s3::{S3Config, S3Storage}; diff --git a/server-new/src/services/storage/engines/local.rs b/server-new/src/services/storage/engines/local.rs index 7a4edbe..9883fa5 100644 --- a/server-new/src/services/storage/engines/local.rs +++ b/server-new/src/services/storage/engines/local.rs @@ -1,10 +1,11 @@ use std::path::{Path, PathBuf}; -use futures::future::BoxFuture; +use futures::{future::BoxFuture, stream::BoxStream}; use tokio::{ fs::File, - io::{AsyncRead, AsyncReadExt, AsyncWriteExt, BufWriter}, + io::{AsyncReadExt, AsyncWriteExt, BufWriter}, }; +use tokio_util::io::StreamReader; use crate::services::storage::{ StorageEngine, @@ -30,7 +31,9 @@ impl StorageEngine for LocalStorage { fn create<'r>( &'r self, file_path: &'r Path, - reader: &'r mut (dyn AsyncRead + Unpin + Send), + _size: usize, + _content_type: &'r str, + stream: BoxStream<'static, Result>, ) -> BoxFuture<'r, StorageResult> { Box::pin(async move { let local_path = self.local_path(file_path); @@ -39,6 +42,7 @@ impl StorageEngine for LocalStorage { let mut file = File::create_new(&local_path).await?; let mut file_writer = BufWriter::new(&mut file); + let mut reader = StreamReader::new(stream); let mut read_buffer = [0; 8192]; let mut total_bytes: usize = 0; diff --git a/server-new/src/services/storage/engines/s3.rs b/server-new/src/services/storage/engines/s3.rs new file mode 100644 index 0000000..0123c37 --- /dev/null +++ b/server-new/src/services/storage/engines/s3.rs @@ -0,0 +1,154 @@ +use std::{ + path::{Path, PathBuf}, + sync::{Arc, atomic::AtomicUsize}, + time::Duration, +}; + +use futures::{TryStreamExt, future::BoxFuture, stream::BoxStream}; +use reqwest::{StatusCode, header}; +use rusty_s3::{Bucket, Credentials, S3Action}; +use serde::{Deserialize, Serialize}; + +use crate::services::storage::{ + StorageEngine, + error::{StorageError, StorageResult}, +}; + +/// Expiration used for S3 requests & presigned URLs +const EXPIRY: Duration = Duration::from_secs(60 * 5); + +/// S3 storage configuration +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct S3Config { + endpoint: reqwest::Url, + bucket: String, + region: String, + access_key: String, + secret_key: String, +} + +/// Storage engine using S3 +pub struct S3Storage<'c> { + base_path: PathBuf, + bucket: Bucket, + credentials: Credentials, + client: &'c reqwest::Client, +} + +impl<'c> S3Storage<'c> { + pub fn new( + base_path: PathBuf, + config: &'c S3Config, + client: &'c reqwest::Client, + ) -> StorageResult { + let bucket = Bucket::new( + config.endpoint.clone(), + rusty_s3::UrlStyle::VirtualHost, + config.bucket.clone(), + config.region.clone(), + ) + .map_err(|e| StorageError::Setup(e.to_string()))?; + let credentials = Credentials::new(&config.access_key, &config.secret_key); + + Ok(Self { + base_path, + bucket, + credentials, + client, + }) + } + + fn file_key(&self, path: &Path) -> Result { + let file_path = self.base_path.join(path); + let file_key = file_path + .to_str() + .ok_or_else(|| StorageError::InvalidPath(file_path.to_string_lossy().into_owned()))?; + + Ok(file_key.to_owned()) + } + + async fn handle_response_error(&self, response: reqwest::Response) -> StorageError { + StorageError::Response(format!( + "Status: {}, Response: {:?}", + response.status().as_u16(), + response.text().await + )) + } +} + +impl StorageEngine for S3Storage<'_> { + fn create<'r>( + &'r self, + path: &'r Path, + size: usize, + content_type: &'r str, + stream: BoxStream<'static, Result>, + ) -> BoxFuture<'r, StorageResult> { + Box::pin(async move { + let file_key = self.file_key(path)?; + let put_object = self.bucket.put_object(Some(&self.credentials), &file_key); + let url = put_object.sign(EXPIRY); + + let total_bytes = Arc::new(AtomicUsize::new(0)); + let size_counter = Arc::clone(&total_bytes); + + let response = self + .client + .put(url) + .header(header::CONTENT_LENGTH, size) + .header(header::CONTENT_TYPE, content_type) + .body(reqwest::Body::wrap_stream(stream.inspect_ok( + move |chunk| { + size_counter.fetch_add(chunk.len(), std::sync::atomic::Ordering::Relaxed); + }, + ))) + .send() + .await?; + + if !response.status().is_success() { + return Err(self.handle_response_error(response).await); + } + + Ok(total_bytes.load(std::sync::atomic::Ordering::Relaxed)) + }) + } + + fn exists<'r>(&'r self, path: &'r Path) -> BoxFuture<'r, StorageResult> { + Box::pin(async move { + let file_key = self.file_key(path)?; + let head_object = self.bucket.head_object(Some(&self.credentials), &file_key); + let url = head_object.sign(EXPIRY); + + let response = self.client.head(url).send().await?; + match response.status() { + StatusCode::OK => Ok(true), + StatusCode::NOT_FOUND => Ok(false), + _ => Err(self.handle_response_error(response).await), + } + }) + } + + fn delete<'r>(&'r self, path: &'r Path) -> BoxFuture<'r, StorageResult<()>> { + Box::pin(async move { + let file_key = self.file_key(path)?; + let delete_object = self + .bucket + .delete_object(Some(&self.credentials), &file_key); + let url = delete_object.sign(EXPIRY); + + let response = self.client.delete(url).send().await?; + match response.status() { + StatusCode::NO_CONTENT => Ok(()), + _ => Err(self.handle_response_error(response).await), + } + }) + } + + fn signed_url(&self, path: &Path) -> StorageResult> { + let file_key = self.file_key(path)?; + let get_object = self.bucket.get_object(Some(&self.credentials), &file_key); + let url = get_object.sign(EXPIRY); + + Ok(Some(url.into())) + } +} diff --git a/server-new/src/services/storage/error.rs b/server-new/src/services/storage/error.rs index bba539a..caad853 100644 --- a/server-new/src/services/storage/error.rs +++ b/server-new/src/services/storage/error.rs @@ -10,15 +10,21 @@ pub enum StorageError { NotFound, #[error("File with this path already exists")] AlreadyExists, - #[error("File is missing content type")] - MissingContentType, #[error("Unsupported content type: '{0}'")] UnsupportedContentType(String), #[error("Invalid file name/path: '{0}'")] InvalidPath(String), + #[error("File had unexpected size: {0} bytes")] + WrongSize(usize), + #[error("Storage setup error: {0}")] + Setup(String), #[error("IO error: {0}")] Io(#[from] std::io::Error), + #[error("Storage request error: {0}")] + Request(#[from] reqwest::Error), + #[error("Storage response error: {0}")] + Response(String), #[error("Database error: {0}")] Database(#[from] diesel::result::Error), #[error(transparent)] @@ -30,7 +36,6 @@ impl From for AppError { match error { StorageError::NotFound => AppError::not_found(error.to_string()), StorageError::AlreadyExists - | StorageError::MissingContentType | StorageError::UnsupportedContentType(_) | StorageError::InvalidPath(_) => Self::bad_request(error.to_string()), err => Self::internal(err.into()), diff --git a/server-new/src/services/storage/interface.rs b/server-new/src/services/storage/interface.rs index 007c3dd..868b7ab 100644 --- a/server-new/src/services/storage/interface.rs +++ b/server-new/src/services/storage/interface.rs @@ -2,8 +2,7 @@ use std::path::Path; -use futures::future::BoxFuture; -use tokio::io::AsyncRead; +use futures::{future::BoxFuture, stream::BoxStream}; use super::error::StorageResult; @@ -12,7 +11,9 @@ pub trait StorageEngine: Send + Sync { fn create<'r>( &'r self, path: &'r Path, - reader: &'r mut (dyn AsyncRead + Unpin + Send), + size: usize, + content_type: &'r str, + stream: BoxStream<'static, Result>, ) -> BoxFuture<'r, StorageResult>; fn exists<'r>(&'r self, path: &'r Path) -> BoxFuture<'r, StorageResult>; fn delete<'r>(&'r self, path: &'r Path) -> BoxFuture<'r, StorageResult<()>>; diff --git a/server-new/src/services/storage/mod.rs b/server-new/src/services/storage/mod.rs index 59d471e..4380151 100644 --- a/server-new/src/services/storage/mod.rs +++ b/server-new/src/services/storage/mod.rs @@ -1,21 +1,21 @@ use std::path::{Path, PathBuf}; use futures::Stream; -use tokio_util::io::StreamReader; use uuid::Uuid; use crate::{ + config::StorageConfig, db::{ DbPool, DbService, models::{ChatRsFile, ChatRsFileType, NewChatRsFile}, }, services::storage::{ - engines::LocalStorage, + engines::{LocalStorage, S3Storage}, error::{StorageError, StorageResult}, }, }; -mod engines; +pub mod engines; mod error; mod interface; @@ -31,11 +31,23 @@ pub const SESSION_FOLDER: &str = "session"; pub struct StorageService<'r> { data_dir: &'r Path, db_pool: &'r DbPool, + config: &'r StorageConfig, + http_client: &'r reqwest::Client, } impl<'r> StorageService<'r> { - pub fn new(data_dir: &'r Path, db_pool: &'r DbPool) -> Self { - Self { data_dir, db_pool } + pub fn new( + data_dir: &'r Path, + db_pool: &'r DbPool, + http_client: &'r reqwest::Client, + config: &'r StorageConfig, + ) -> Self { + Self { + data_dir, + db_pool, + config, + http_client, + } } pub async fn create_file( @@ -43,21 +55,27 @@ impl<'r> StorageService<'r> { user_id: &Uuid, session_id: Option<&Uuid>, path: &str, - content_type: Option<&str>, - stream: impl Stream> + Send + Unpin, + size: usize, + content_type: &str, + stream: impl Stream> + Send + 'static, ) -> StorageResult { let file_path = self.build_file_path(user_id, session_id, &path)?; - let content_type = content_type.ok_or(StorageError::MissingContentType)?; let file_type = self.validate_file_type(content_type)?; - let storage = self.storage_engine(); + let storage = self.storage_engine()?; if storage.exists(&file_path).await? { return Err(StorageError::AlreadyExists); } - let mut reader = StreamReader::new(stream); - let file_size = match storage.create(&file_path, &mut reader).await { - Ok(num_bytes) => num_bytes, + match storage + .create(&file_path, size, content_type, Box::pin(stream)) + .await + { + Ok(bytes_written) if bytes_written == size => {} + Ok(wrong_size) => { + let _ = storage.delete(&file_path).await; + return Err(StorageError::WrongSize(wrong_size)); + } Err(err) => { let _ = storage.delete(&file_path).await; return Err(err); @@ -71,7 +89,7 @@ impl<'r> StorageService<'r> { path, file_type: file_type.as_ref(), content_type, - size: file_size.try_into().unwrap_or_default(), + size: size.try_into().unwrap_or_default(), }; let db_file = db.files().create_file(new_file).await?; @@ -95,7 +113,7 @@ impl<'r> StorageService<'r> { } .ok_or(StorageError::NotFound)?; - let storage = self.storage_engine(); + let storage = self.storage_engine()?; let file_path = self.build_file_path(user_id, session_id, &db_file.path)?; if let Err(err) = storage.delete(&file_path).await { tracing::warn!("error deleting file {file_id} with path {file_path:?}: {err}"); @@ -113,8 +131,15 @@ impl<'r> StorageService<'r> { Ok(deleted_file_id) } - fn storage_engine(&self) -> Box { - Box::new(LocalStorage::new(self.data_dir.join(STORAGE_FOLDER))) + fn storage_engine(&self) -> StorageResult> { + Ok(match self.config { + StorageConfig::Local => Box::new(LocalStorage::new(self.data_dir.join(STORAGE_FOLDER))), + StorageConfig::S3(config) => Box::new(S3Storage::new( + STORAGE_FOLDER.into(), + config, + self.http_client, + )?), + }) } fn build_file_path( @@ -149,6 +174,7 @@ impl<'r> StorageService<'r> { match content_type { "image/jpeg" | "image/png" | "image/webp" => Ok(ChatRsFileType::Image), "application/pdf" => Ok(ChatRsFileType::Pdf), + "application/json" | "application/xml" => Ok(ChatRsFileType::Text), text if text.starts_with("text/") => Ok(ChatRsFileType::Text), unsupported => Err(StorageError::UnsupportedContentType(unsupported.into())), } diff --git a/server-new/src/state.rs b/server-new/src/state.rs index a8baf1a..223306a 100644 --- a/server-new/src/state.rs +++ b/server-new/src/state.rs @@ -46,7 +46,12 @@ impl AppState { ModelService::new(&self.redis, &self.http_client) } pub fn storage_service(&self) -> StorageService<'_> { - StorageService::new(&self.config.server.data_dir, &self.db_pool) + StorageService::new( + &self.config.server.data_dir, + &self.db_pool, + &self.http_client, + &self.config.storage, + ) } }