Skip to main content

systemprompt_api/services/gateway/service/credentials/
google.rs

1//! Google service-account credentials: RS256 assertion in, access token out.
2//!
3//! Implements the JWT-bearer profile (RFC 7523) that Google's token endpoint
4//! accepts: sign a short assertion with the service account's private key,
5//! POST it as `urn:ietf:params:oauth:grant-type:jwt-bearer`, receive an access
6//! token valid for about an hour.
7//!
8//! Tokens are cached per secret name and reused until shortly before they
9//! expire, because minting on every request would add a round trip to Google
10//! in front of every round trip to the model. The skew is deliberate: a token
11//! that expires in flight fails the *user's* request, so it is retired early
12//! rather than used to the last second.
13//!
14//! Copyright (c) systemprompt.io — Business Source License 1.1.
15//! See <https://systemprompt.io> for licensing details.
16
17use std::collections::HashMap;
18use std::sync::{OnceLock, PoisonError, RwLock};
19use std::time::{Duration, SystemTime, UNIX_EPOCH};
20
21use anyhow::{Result, anyhow, bail};
22use jsonwebtoken::{Algorithm, EncodingKey, Header};
23use serde::{Deserialize, Serialize};
24
25// Why: Google limits a service-account JWT assertion's lifetime to one hour.
26const ASSERTION_TTL: Duration = Duration::from_secs(3600);
27
28const EXPIRY_SKEW: Duration = Duration::from_secs(120);
29
30const SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform";
31
32#[derive(Debug, Deserialize)]
33pub struct ServiceAccountKey {
34    pub(super) client_email: String,
35    pub(super) private_key: String,
36    #[serde(default = "default_token_uri")]
37    pub token_uri: String,
38}
39
40fn default_token_uri() -> String {
41    "https://oauth2.googleapis.com/token".to_owned()
42}
43
44impl ServiceAccountKey {
45    pub fn parse(secret: &str) -> Result<Option<Self>> {
46        let Ok(value) = serde_json::from_str::<serde_json::Value>(secret) else {
47            return Ok(None);
48        };
49        if value.get("type").and_then(serde_json::Value::as_str) != Some("service_account") {
50            return Ok(None);
51        }
52        serde_json::from_value(value)
53            .map(Some)
54            .map_err(|e| anyhow!("service-account key is malformed: {e}"))
55    }
56}
57
58#[derive(Debug, Serialize)]
59struct Assertion<'a> {
60    iss: &'a str,
61    scope: &'a str,
62    aud: &'a str,
63    iat: u64,
64    exp: u64,
65}
66
67#[derive(Debug, Deserialize)]
68struct TokenResponse {
69    access_token: String,
70    #[serde(default)]
71    expires_in: u64,
72}
73
74#[derive(Debug, Clone)]
75struct CachedToken {
76    token: String,
77    expires_at: SystemTime,
78}
79
80fn cache() -> &'static RwLock<HashMap<String, CachedToken>> {
81    static CACHE: OnceLock<RwLock<HashMap<String, CachedToken>>> = OnceLock::new();
82    CACHE.get_or_init(|| RwLock::new(HashMap::new()))
83}
84
85pub async fn access_token(secret_name: &str, key: &ServiceAccountKey) -> Result<String> {
86    if let Some(token) = cached(secret_name) {
87        return Ok(token);
88    }
89
90    let response = exchange(key).await?;
91    let ttl = if response.expires_in == 0 {
92        ASSERTION_TTL
93    } else {
94        Duration::from_secs(response.expires_in)
95    };
96
97    cache()
98        .write()
99        .unwrap_or_else(PoisonError::into_inner)
100        .insert(
101            secret_name.to_owned(),
102            CachedToken {
103                token: response.access_token.clone(),
104                expires_at: SystemTime::now() + ttl,
105            },
106        );
107
108    Ok(response.access_token)
109}
110
111fn cached(secret_name: &str) -> Option<String> {
112    let guard = cache().read().unwrap_or_else(PoisonError::into_inner);
113    let token = guard.get(secret_name).and_then(|entry| {
114        (entry.expires_at > SystemTime::now() + EXPIRY_SKEW).then(|| entry.token.clone())
115    });
116    drop(guard);
117    token
118}
119
120async fn exchange(key: &ServiceAccountKey) -> Result<TokenResponse> {
121    let assertion = sign_assertion(key)?;
122
123    let response = reqwest::Client::new()
124        .post(&key.token_uri)
125        .form(&[
126            (
127                "grant_type",
128                "urn:ietf:params:oauth:grant-type:jwt-bearer".to_owned(),
129            ),
130            ("assertion", assertion),
131        ])
132        .send()
133        .await
134        .map_err(|e| anyhow!("token endpoint {} unreachable: {e}", key.token_uri))?;
135
136    let status = response.status();
137    let body = response.text().await.unwrap_or_default();
138    if !status.is_success() {
139        bail!("token endpoint returned {status}: {}", body.trim());
140    }
141
142    serde_json::from_str(&body)
143        .map_err(|e| anyhow!("token endpoint returned an unreadable body: {e}"))
144}
145
146fn sign_assertion(key: &ServiceAccountKey) -> Result<String> {
147    let now = SystemTime::now()
148        .duration_since(UNIX_EPOCH)
149        .map_err(|e| anyhow!("system clock is before the unix epoch: {e}"))?
150        .as_secs();
151
152    let claims = Assertion {
153        iss: &key.client_email,
154        scope: SCOPE,
155        aud: &key.token_uri,
156        iat: now,
157        exp: now + ASSERTION_TTL.as_secs(),
158    };
159
160    let encoding = EncodingKey::from_rsa_pem(key.private_key.as_bytes())
161        .map_err(|e| anyhow!("service-account private_key is not a valid RSA PEM: {e}"))?;
162
163    jsonwebtoken::encode(&Header::new(Algorithm::RS256), &claims, &encoding)
164        .map_err(|e| anyhow!("could not sign the assertion: {e}"))
165}