Skip to main content

systemprompt_users/sessions/
ai_provider.rs

1//! AI session lifecycle and usage accounting owned by users.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use async_trait::async_trait;
7use systemprompt_identifiers::SessionId;
8use systemprompt_traits::{
9    ActiveSession, AiProviderError, AiProviderResult, AiSessionProvider, CreateAiSessionParams,
10};
11
12use super::SessionRepository;
13use systemprompt_traits::session_store::CreateSessionParams;
14
15#[derive(Debug)]
16pub struct UsersAiSessionProvider {
17    session_repo: SessionRepository,
18}
19
20impl UsersAiSessionProvider {
21    pub const fn from_repository(session_repo: SessionRepository) -> Self {
22        Self { session_repo }
23    }
24}
25
26#[async_trait]
27impl AiSessionProvider for UsersAiSessionProvider {
28    async fn create_session(&self, params: CreateAiSessionParams<'_>) -> AiProviderResult<()> {
29        let full_params = CreateSessionParams {
30            session_id: params.session_id,
31            user_id: params.user_id,
32            session_source: params.session_source,
33            fingerprint_hash: None,
34            ip_address: None,
35            user_agent: None,
36            device_type: None,
37            browser: None,
38            os: None,
39            country: None,
40            region: None,
41            city: None,
42            preferred_locale: None,
43            referrer_source: None,
44            referrer_url: None,
45            landing_page: None,
46            entry_url: None,
47            utm_source: None,
48            utm_medium: None,
49            utm_campaign: None,
50            utm_content: None,
51            utm_term: None,
52            is_bot: false,
53            is_ai_crawler: false,
54            expires_at: params.expires_at,
55        };
56
57        self.session_repo
58            .create_session(&full_params)
59            .await
60            .map_err(|e| AiProviderError::Internal(e.to_string()))
61    }
62
63    async fn increment_ai_usage(
64        &self,
65        session_id: &SessionId,
66        tokens: i32,
67        cost_microdollars: i64,
68    ) -> AiProviderResult<()> {
69        self.session_repo
70            .increment_ai_usage(session_id, tokens, cost_microdollars)
71            .await
72            .map_err(|e| AiProviderError::Internal(e.to_string()))
73    }
74
75    async fn find_live_session(
76        &self,
77        session_id: &SessionId,
78    ) -> AiProviderResult<Option<ActiveSession>> {
79        let session = self
80            .session_repo
81            .find_active_by_id(session_id)
82            .await
83            .map_err(|e| AiProviderError::Internal(e.to_string()))?;
84        Ok(session.map(|row| ActiveSession {
85            user_id: row.user_id,
86        }))
87    }
88}