systemprompt_api/services/gateway/service/credentials/
google.rs1use 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
25const 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}