1use base64::Engine;
2use chrono::Utc;
3use sea_orm::entity::prelude::*;
4use sea_orm::{DatabaseConnection, Set};
5use sha2::{Digest, Sha256};
6
7use crate::auth::error::AuthError;
8use crate::auth::principal::{AuthPrincipal, PrincipalSource};
9
10#[derive(Clone, Debug, PartialEq, DeriveEntityModel)]
11#[sea_orm(table_name = "miryad_api_tokens")]
12pub struct Model {
13 #[sea_orm(primary_key)]
14 pub id: i32,
15 pub subject: String,
18 pub name: String,
20 pub token_hash: String,
22 pub created_at: DateTimeUtc,
23 pub expires_at: Option<DateTimeUtc>,
24 pub last_used_at: Option<DateTimeUtc>,
25}
26
27#[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)]
28pub enum Relation {}
29
30impl ActiveModelBehavior for ActiveModel {}
31
32pub type ApiToken = Entity;
34
35pub struct IssuedToken {
38 pub id: i32,
39 pub token: String,
40}
41
42fn generate_token() -> String {
43 use rand::RngCore;
44 let mut bytes = [0u8; 32];
45 rand::thread_rng().fill_bytes(&mut bytes);
46 format!(
47 "mrd_{}",
48 base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes)
49 )
50}
51
52fn hash_token(token: &str) -> String {
53 hex::encode(Sha256::digest(token.as_bytes()))
54}
55
56pub async fn issue_token(
57 db: &DatabaseConnection,
58 subject: &str,
59 name: &str,
60 expires_at: Option<DateTimeUtc>,
61) -> Result<IssuedToken, AuthError> {
62 let token = generate_token();
63 let active = ActiveModel {
64 subject: Set(subject.to_string()),
65 name: Set(name.to_string()),
66 token_hash: Set(hash_token(&token)),
67 created_at: Set(Utc::now()),
68 expires_at: Set(expires_at),
69 last_used_at: Set(None),
70 ..Default::default()
71 };
72 let inserted = active.insert(db).await?;
73
74 Ok(IssuedToken {
75 id: inserted.id,
76 token,
77 })
78}
79
80pub async fn validate_token(db: &DatabaseConnection, token: &str) -> Result<AuthPrincipal, AuthError> {
81 let record = Entity::find()
82 .filter(Column::TokenHash.eq(hash_token(token)))
83 .one(db)
84 .await?
85 .ok_or(AuthError::InvalidToken)?;
86
87 if let Some(expires_at) = record.expires_at
88 && expires_at <= Utc::now()
89 {
90 return Err(AuthError::TokenExpired);
91 }
92
93 let id = record.id;
94 let subject = record.subject.clone();
95 let mut active: ActiveModel = record.into();
96 active.last_used_at = Set(Some(Utc::now()));
97 active.update(db).await?;
98
99 Ok(AuthPrincipal {
100 subject,
101 email: None,
102 source: PrincipalSource::ApiToken { token_id: id },
103 })
104}
105
106pub async fn revoke_token(db: &DatabaseConnection, id: i32) -> Result<(), AuthError> {
107 Entity::delete_by_id(id).exec(db).await?;
108 Ok(())
109}
110
111pub async fn ensure_token(
116 db: &DatabaseConnection,
117 subject: &str,
118 name: &str,
119 token: &str,
120 expires_at: Option<DateTimeUtc>,
121) -> Result<(), AuthError> {
122 let hash = hash_token(token);
123 let exists = Entity::find()
124 .filter(Column::TokenHash.eq(&hash))
125 .one(db)
126 .await?
127 .is_some();
128 if exists {
129 return Ok(());
130 }
131
132 let active = ActiveModel {
133 subject: Set(subject.to_string()),
134 name: Set(name.to_string()),
135 token_hash: Set(hash),
136 created_at: Set(Utc::now()),
137 expires_at: Set(expires_at),
138 last_used_at: Set(None),
139 ..Default::default()
140 };
141 active.insert(db).await?;
142 Ok(())
143}
144
145#[cfg(test)]
146mod tests {
147 use super::*;
148 use crate::migration::Migrator;
149 use sea_orm_migration::MigratorTrait;
150
151 async fn test_db() -> DatabaseConnection {
152 let db = sea_orm::Database::connect("sqlite::memory:")
153 .await
154 .expect("in-memory sqlite connects");
155 Migrator::up(&db, None).await.expect("migrations apply cleanly");
156 db
157 }
158
159 #[tokio::test]
160 async fn issued_token_validates_and_updates_last_used() {
161 let db = test_db().await;
162 let issued = issue_token(&db, "user-123", "test token", None)
163 .await
164 .expect("issuing succeeds");
165
166 assert_ne!(issued.token, hash_token(&issued.token));
167
168 let principal = validate_token(&db, &issued.token).await.expect("token is valid");
169 assert_eq!(principal.subject, "user-123");
170 assert!(matches!(
171 principal.source,
172 PrincipalSource::ApiToken { token_id } if token_id == issued.id
173 ));
174
175 let record = Entity::find_by_id(issued.id)
176 .one(&db)
177 .await
178 .expect("query succeeds")
179 .expect("record exists");
180 assert!(record.last_used_at.is_some());
181 }
182
183 #[tokio::test]
184 async fn expired_token_is_rejected() {
185 let db = test_db().await;
186 let issued = issue_token(&db, "user-123", "expired token", Some(Utc::now()))
187 .await
188 .expect("issuing succeeds");
189
190 let result = validate_token(&db, &issued.token).await;
191 assert!(matches!(result, Err(AuthError::TokenExpired)));
192 }
193
194 #[tokio::test]
195 async fn unknown_token_is_rejected() {
196 let db = test_db().await;
197 let result = validate_token(&db, "mrd_does-not-exist").await;
198 assert!(matches!(result, Err(AuthError::InvalidToken)));
199 }
200
201 #[tokio::test]
202 async fn revoked_token_is_rejected() {
203 let db = test_db().await;
204 let issued = issue_token(&db, "user-123", "to revoke", None)
205 .await
206 .expect("issuing succeeds");
207 revoke_token(&db, issued.id).await.expect("revocation succeeds");
208
209 let result = validate_token(&db, &issued.token).await;
210 assert!(matches!(result, Err(AuthError::InvalidToken)));
211 }
212
213 #[tokio::test]
214 async fn ensure_token_creates_then_authenticates() {
215 let db = test_db().await;
216 ensure_token(&db, "service-account", "bootstrap", "mrd_fixed-secret", None)
217 .await
218 .expect("ensure succeeds");
219
220 let principal = validate_token(&db, "mrd_fixed-secret")
221 .await
222 .expect("token is valid");
223 assert_eq!(principal.subject, "service-account");
224 }
225
226 #[tokio::test]
227 async fn ensure_token_is_idempotent() {
228 let db = test_db().await;
229 ensure_token(&db, "service-account", "bootstrap", "mrd_fixed-secret", None)
230 .await
231 .expect("first ensure succeeds");
232 ensure_token(&db, "service-account", "bootstrap", "mrd_fixed-secret", None)
233 .await
234 .expect("second ensure succeeds");
235
236 let count = Entity::find()
237 .filter(Column::TokenHash.eq(hash_token("mrd_fixed-secret")))
238 .all(&db)
239 .await
240 .expect("query succeeds")
241 .len();
242 assert_eq!(count, 1);
243 }
244}