diff --git a/src/application/routes/roasters.rs b/src/application/routes/roasters.rs index 2205c4c..5b1d372 100644 --- a/src/application/routes/roasters.rs +++ b/src/application/routes/roasters.rs @@ -118,15 +118,15 @@ define_delete_handler!( render_roaster_list_fragment ); -#[tracing::instrument(skip(state, _auth_user, headers, payload))] +#[tracing::instrument(skip(state, auth_user, headers, payload))] pub(crate) async fn extract_roaster( State(state): State, - _auth_user: AuthenticatedUser, + auth_user: AuthenticatedUser, headers: HeaderMap, payload: FlexiblePayload, ) -> Result { let (input, _) = payload.into_parts(); - let result = ai::extract_roaster( + let (result, usage) = ai::extract_roaster( &state.http_client, &state.openrouter_api_key, &state.openrouter_model, @@ -135,6 +135,14 @@ pub(crate) async fn extract_roaster( .await .map_err(ApiError::from)?; + super::support::record_ai_usage( + state.ai_usage_repo.clone(), + auth_user.0.id, + &state.openrouter_model, + "extract-roaster", + usage, + ); + if is_datastar_request(&headers) { use serde_json::Value; let signals = vec![ diff --git a/src/application/routes/roasts.rs b/src/application/routes/roasts.rs index 692f249..422a800 100644 --- a/src/application/routes/roasts.rs +++ b/src/application/routes/roasts.rs @@ -257,15 +257,15 @@ impl TastingNotesInput { } } -#[tracing::instrument(skip(state, _auth_user, headers, payload))] +#[tracing::instrument(skip(state, auth_user, headers, payload))] pub(crate) async fn extract_roast_info( State(state): State, - _auth_user: AuthenticatedUser, + auth_user: AuthenticatedUser, headers: HeaderMap, payload: FlexiblePayload, ) -> Result { let (input, _) = payload.into_parts(); - let result = ai::extract_roast( + let (result, usage) = ai::extract_roast( &state.http_client, &state.openrouter_api_key, &state.openrouter_model, @@ -274,6 +274,14 @@ pub(crate) async fn extract_roast_info( .await .map_err(ApiError::from)?; + super::support::record_ai_usage( + state.ai_usage_repo.clone(), + auth_user.0.id, + &state.openrouter_model, + "extract-roast", + usage, + ); + if is_datastar_request(&headers) { use serde_json::Value; diff --git a/src/application/routes/scan.rs b/src/application/routes/scan.rs index 16c7c1b..f3b92d7 100644 --- a/src/application/routes/scan.rs +++ b/src/application/routes/scan.rs @@ -14,17 +14,17 @@ use crate::domain::errors::RepositoryError; use crate::domain::roasters::NewRoaster; use crate::domain::roasts::NewRoast; use crate::domain::timeline::{NewTimelineEvent, TimelineEventDetail}; -use crate::infrastructure::ai::{self, ExtractionInput}; +use crate::infrastructure::ai::{self, ExtractionInput, Usage}; -#[tracing::instrument(skip(state, _auth_user, headers, payload))] +#[tracing::instrument(skip(state, auth_user, headers, payload))] pub(crate) async fn extract_bag_scan( State(state): State, - _auth_user: AuthenticatedUser, + auth_user: AuthenticatedUser, headers: HeaderMap, payload: FlexiblePayload, ) -> Result { let (input, _) = payload.into_parts(); - let result = ai::extract_bag_scan( + let (result, usage) = ai::extract_bag_scan( &state.http_client, &state.openrouter_api_key, &state.openrouter_model, @@ -33,6 +33,14 @@ pub(crate) async fn extract_bag_scan( .await .map_err(ApiError::from)?; + super::support::record_ai_usage( + state.ai_usage_repo.clone(), + auth_user.0.id, + &state.openrouter_model, + "extract-bag-scan", + usage, + ); + if is_datastar_request(&headers) { use serde_json::Value; @@ -130,15 +138,16 @@ struct ScanResult { } /// Populate a `BagScanSubmission` from AI extraction when image/prompt is provided. +/// Returns the usage data so the caller can record it. async fn extract_into_submission( state: &AppState, submission: &mut BagScanSubmission, -) -> Result<(), ApiError> { +) -> Result, ApiError> { let input = ExtractionInput { image: submission.image.take(), prompt: submission.prompt.take(), }; - let result = ai::extract_bag_scan( + let (result, usage) = ai::extract_bag_scan( &state.http_client, &state.openrouter_api_key, &state.openrouter_model, @@ -178,14 +187,14 @@ async fn extract_into_submission( submission.tasting_notes = TastingNotesInput::Text(notes.join(", ")); } - Ok(()) + Ok(usage) } #[allow(clippy::too_many_lines)] -#[tracing::instrument(skip(state, _auth_user, headers, payload))] +#[tracing::instrument(skip(state, auth_user, headers, payload))] pub(crate) async fn submit_scan( State(state): State, - _auth_user: AuthenticatedUser, + auth_user: AuthenticatedUser, headers: HeaderMap, payload: FlexiblePayload, ) -> Result { @@ -196,7 +205,14 @@ pub(crate) async fn submit_scan( || submission.prompt.as_deref().is_some_and(|s| !s.is_empty()); if has_raw_input { - extract_into_submission(&state, &mut submission).await?; + let usage = extract_into_submission(&state, &mut submission).await?; + super::support::record_ai_usage( + state.ai_usage_repo.clone(), + auth_user.0.id, + &state.openrouter_model, + "extract-bag-scan", + usage, + ); } // Build and normalize the roaster diff --git a/src/application/routes/support.rs b/src/application/routes/support.rs index 6136da8..26afe33 100644 --- a/src/application/routes/support.rs +++ b/src/application/routes/support.rs @@ -236,6 +236,31 @@ pub(super) async fn load_cafe_options(state: &AppState) -> Result, + user_id: crate::domain::ids::UserId, + model: &str, + endpoint: &str, + usage: Option, +) { + let Some(usage) = usage else { return }; + let new_usage = crate::domain::ai_usage::NewAiUsage { + user_id, + model: model.to_string(), + endpoint: endpoint.to_string(), + prompt_tokens: usage.prompt_tokens, + completion_tokens: usage.completion_tokens, + total_tokens: usage.total_tokens, + cost: usage.cost, + }; + tokio::spawn(async move { + if let Err(err) = repo.insert(new_usage).await { + tracing::warn!(error = %err, "failed to record AI usage"); + } + }); +} + pub fn is_datastar_request(headers: &HeaderMap) -> bool { headers .get("datastar-request") diff --git a/src/application/server.rs b/src/application/server.rs index bb22432..ba4e0b9 100644 --- a/src/application/server.rs +++ b/src/application/server.rs @@ -12,13 +12,14 @@ use webauthn_rs::prelude::*; use crate::application::routes::app_router; use crate::domain::registration_tokens::NewRegistrationToken; use crate::domain::repositories::{ - BagRepository, BrewRepository, CafeRepository, CupRepository, GearRepository, - PasskeyCredentialRepository, RegistrationTokenRepository, RoastRepository, RoasterRepository, - SessionRepository, TimelineEventRepository, TokenRepository, UserRepository, + AiUsageRepository, BagRepository, BrewRepository, CafeRepository, CupRepository, + GearRepository, PasskeyCredentialRepository, RegistrationTokenRepository, RoastRepository, + RoasterRepository, SessionRepository, TimelineEventRepository, TokenRepository, UserRepository, }; use crate::infrastructure::auth::{generate_session_token, hash_token}; use crate::infrastructure::backup::BackupService; use crate::infrastructure::database::Database; +use crate::infrastructure::repositories::ai_usage::SqlAiUsageRepository; use crate::infrastructure::repositories::bags::SqlBagRepository; use crate::infrastructure::repositories::brews::SqlBrewRepository; use crate::infrastructure::repositories::cafes::SqlCafeRepository; @@ -59,6 +60,7 @@ pub struct AppState { pub session_repo: Arc, pub passkey_repo: Arc, pub registration_token_repo: Arc, + pub ai_usage_repo: Arc, pub webauthn: Arc, pub challenge_store: Arc, pub http_client: reqwest::Client, @@ -101,6 +103,8 @@ pub async fn serve(config: ServerConfig) -> anyhow::Result<()> { Arc::new(SqlPasskeyCredentialRepository::new(database.clone_pool())); let registration_token_repo: Arc = Arc::new(SqlRegistrationTokenRepository::new(database.clone_pool())); + let ai_usage_repo: Arc = + Arc::new(SqlAiUsageRepository::new(database.clone_pool())); let backup_service = Arc::new(BackupService::new(database.clone_pool())); let challenge_store = Arc::new(ChallengeStore::new()); @@ -122,6 +126,7 @@ pub async fn serve(config: ServerConfig) -> anyhow::Result<()> { session_repo, passkey_repo, registration_token_repo, + ai_usage_repo, webauthn, challenge_store, http_client: reqwest::Client::new(), diff --git a/tests/server/helpers.rs b/tests/server/helpers.rs index 427b0ee..692a47f 100644 --- a/tests/server/helpers.rs +++ b/tests/server/helpers.rs @@ -164,6 +164,11 @@ async fn spawn_app_inner( session_repo, passkey_repo, registration_token_repo, + ai_usage_repo: Arc::new( + brewlog::infrastructure::repositories::ai_usage::SqlAiUsageRepository::new( + _database.clone_pool(), + ), + ), webauthn: test_webauthn(), challenge_store: Arc::new(ChallengeStore::new()), http_client: reqwest::Client::new(),