1use std::collections::HashSet;
4
5use chrono::{DateTime, NaiveTime, Utc};
6use rand::Rng;
7use serde::{Deserialize, Deserializer, Serialize, Serializer, de};
8use sha2::{Digest, Sha256};
9
10use crate::limits::Limits;
11use crate::scope::Scope;
12
13#[derive(Debug, Clone)]
15pub struct KeyRecord {
16 pub id: String,
18 pub hash: String,
20 pub scopes: HashSet<Scope>,
21
22 pub allowed_markets: Option<HashSet<String>>,
23 pub allowed_symbols: Option<HashSet<String>>,
24 pub max_order_value: Option<f64>,
25 pub max_daily_value: Option<f64>,
26 pub hours_window: Option<String>,
28 pub max_orders_per_minute: Option<u32>,
30 pub allowed_trd_sides: Option<HashSet<String>>,
32 pub allowed_acc_ids: Option<HashSet<u64>>,
35 pub allowed_card_nums: Option<Vec<String>>,
54 pub expires_at: Option<DateTime<Utc>>,
55 pub created_at: DateTime<Utc>,
56 pub note: Option<String>,
57 pub allowed_machines: Option<Vec<String>>,
66
67 pub raw_explicit_acc_ids: Option<HashSet<u64>>,
81}
82
83#[derive(Debug, Default)]
84enum FieldPresence<T> {
85 #[default]
86 Missing,
87 Present(Option<T>),
88}
89
90impl<'de, T> Deserialize<'de> for FieldPresence<T>
91where
92 T: Deserialize<'de>,
93{
94 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
95 where
96 D: Deserializer<'de>,
97 {
98 Option::<T>::deserialize(deserializer).map(Self::Present)
99 }
100}
101
102#[derive(Debug, Default, Deserialize)]
103#[serde(deny_unknown_fields)]
104struct KeyLimitsWire {
105 #[serde(default)]
106 allowed_markets: FieldPresence<HashSet<String>>,
107 #[serde(default)]
108 allowed_symbols: FieldPresence<HashSet<String>>,
109 #[serde(default)]
110 max_order_value: FieldPresence<f64>,
111 #[serde(default)]
112 max_daily_value: FieldPresence<f64>,
113 #[serde(default)]
114 hours_window: FieldPresence<String>,
115 #[serde(default)]
116 max_orders_per_minute: FieldPresence<u32>,
117 #[serde(default)]
118 allowed_trd_sides: FieldPresence<HashSet<String>>,
119 #[serde(default)]
120 allowed_acc_ids: FieldPresence<HashSet<u64>>,
121 #[serde(default)]
122 allowed_card_nums: FieldPresence<Vec<String>>,
123}
124
125#[derive(Debug, Deserialize)]
126#[serde(deny_unknown_fields)]
127struct KeyRecordWire {
128 id: String,
129 hash: String,
130 scopes: HashSet<Scope>,
131 #[serde(default)]
132 limits: FieldPresence<KeyLimitsWire>,
133 #[serde(default)]
134 allowed_markets: FieldPresence<HashSet<String>>,
135 #[serde(default)]
136 allowed_symbols: FieldPresence<HashSet<String>>,
137 #[serde(default)]
138 max_order_value: FieldPresence<f64>,
139 #[serde(default)]
140 max_daily_value: FieldPresence<f64>,
141 #[serde(default)]
142 hours_window: FieldPresence<String>,
143 #[serde(default)]
144 max_orders_per_minute: FieldPresence<u32>,
145 #[serde(default)]
146 allowed_trd_sides: FieldPresence<HashSet<String>>,
147 #[serde(default)]
148 allowed_acc_ids: FieldPresence<HashSet<u64>>,
149 #[serde(default)]
150 allowed_card_nums: FieldPresence<Vec<String>>,
151 #[serde(default)]
152 expires_at: Option<DateTime<Utc>>,
153 created_at: DateTime<Utc>,
154 #[serde(default)]
155 note: Option<String>,
156 #[serde(default)]
157 allowed_machines: Option<Vec<String>>,
158}
159
160fn merge_limit_field<T, E>(
161 name: &'static str,
162 flat: FieldPresence<T>,
163 nested: FieldPresence<T>,
164) -> Result<Option<T>, E>
165where
166 E: de::Error,
167{
168 match (flat, nested) {
169 (FieldPresence::Present(_), FieldPresence::Present(_)) => Err(E::custom(format!(
170 "duplicate limit field {name:?}: configure it either at the key root (legacy) or in limits, not both"
171 ))),
172 (FieldPresence::Present(value), FieldPresence::Missing)
173 | (FieldPresence::Missing, FieldPresence::Present(value)) => Ok(value),
174 (FieldPresence::Missing, FieldPresence::Missing) => Ok(None),
175 }
176}
177
178impl<'de> Deserialize<'de> for KeyRecord {
179 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
180 where
181 D: Deserializer<'de>,
182 {
183 let wire = KeyRecordWire::deserialize(deserializer)?;
184 let nested = match wire.limits {
185 FieldPresence::Missing | FieldPresence::Present(None) => KeyLimitsWire::default(),
186 FieldPresence::Present(Some(limits)) => limits,
187 };
188 let allowed_markets = merge_limit_field::<_, D::Error>(
189 "allowed_markets",
190 wire.allowed_markets,
191 nested.allowed_markets,
192 )?;
193 let allowed_symbols = merge_limit_field::<_, D::Error>(
194 "allowed_symbols",
195 wire.allowed_symbols,
196 nested.allowed_symbols,
197 )?;
198 let max_order_value = merge_limit_field::<_, D::Error>(
199 "max_order_value",
200 wire.max_order_value,
201 nested.max_order_value,
202 )?;
203 let max_daily_value = merge_limit_field::<_, D::Error>(
204 "max_daily_value",
205 wire.max_daily_value,
206 nested.max_daily_value,
207 )?;
208 let hours_window = merge_limit_field::<_, D::Error>(
209 "hours_window",
210 wire.hours_window,
211 nested.hours_window,
212 )?;
213 let max_orders_per_minute = merge_limit_field::<_, D::Error>(
214 "max_orders_per_minute",
215 wire.max_orders_per_minute,
216 nested.max_orders_per_minute,
217 )?;
218 let allowed_trd_sides = merge_limit_field::<_, D::Error>(
219 "allowed_trd_sides",
220 wire.allowed_trd_sides,
221 nested.allowed_trd_sides,
222 )?;
223 let allowed_acc_ids = merge_limit_field::<_, D::Error>(
224 "allowed_acc_ids",
225 wire.allowed_acc_ids,
226 nested.allowed_acc_ids,
227 )?;
228 let allowed_card_nums = merge_limit_field::<_, D::Error>(
229 "allowed_card_nums",
230 wire.allowed_card_nums,
231 nested.allowed_card_nums,
232 )?;
233 let raw_explicit_acc_ids = allowed_acc_ids.clone();
234
235 Ok(Self {
236 id: wire.id,
237 hash: wire.hash,
238 scopes: wire.scopes,
239 allowed_markets,
240 allowed_symbols,
241 max_order_value,
242 max_daily_value,
243 hours_window,
244 max_orders_per_minute,
245 allowed_trd_sides,
246 allowed_acc_ids,
247 allowed_card_nums,
248 expires_at: wire.expires_at,
249 created_at: wire.created_at,
250 note: wire.note,
251 allowed_machines: wire.allowed_machines,
252 raw_explicit_acc_ids,
253 })
254 }
255}
256
257#[derive(Serialize)]
258struct KeyLimitsOut<'a> {
259 #[serde(skip_serializing_if = "Option::is_none")]
260 allowed_markets: Option<&'a HashSet<String>>,
261 #[serde(skip_serializing_if = "Option::is_none")]
262 allowed_symbols: Option<&'a HashSet<String>>,
263 #[serde(skip_serializing_if = "Option::is_none")]
264 max_order_value: Option<f64>,
265 #[serde(skip_serializing_if = "Option::is_none")]
266 max_daily_value: Option<f64>,
267 #[serde(skip_serializing_if = "Option::is_none")]
268 hours_window: Option<&'a String>,
269 #[serde(skip_serializing_if = "Option::is_none")]
270 max_orders_per_minute: Option<u32>,
271 #[serde(skip_serializing_if = "Option::is_none")]
272 allowed_trd_sides: Option<&'a HashSet<String>>,
273 #[serde(skip_serializing_if = "Option::is_none")]
274 allowed_acc_ids: Option<&'a HashSet<u64>>,
275 #[serde(skip_serializing_if = "Option::is_none")]
276 allowed_card_nums: Option<&'a Vec<String>>,
277}
278
279impl KeyLimitsOut<'_> {
280 fn is_empty(&self) -> bool {
281 self.allowed_markets.is_none()
282 && self.allowed_symbols.is_none()
283 && self.max_order_value.is_none()
284 && self.max_daily_value.is_none()
285 && self.hours_window.is_none()
286 && self.max_orders_per_minute.is_none()
287 && self.allowed_trd_sides.is_none()
288 && self.allowed_acc_ids.is_none()
289 && self.allowed_card_nums.is_none()
290 }
291}
292
293#[derive(Serialize)]
294struct KeyRecordOut<'a> {
295 id: &'a str,
296 hash: &'a str,
297 scopes: &'a HashSet<Scope>,
298 #[serde(skip_serializing_if = "Option::is_none")]
299 limits: Option<KeyLimitsOut<'a>>,
300 #[serde(skip_serializing_if = "Option::is_none")]
301 expires_at: Option<DateTime<Utc>>,
302 created_at: DateTime<Utc>,
303 #[serde(skip_serializing_if = "Option::is_none")]
304 note: Option<&'a String>,
305 #[serde(skip_serializing_if = "Option::is_none")]
306 allowed_machines: Option<&'a Vec<String>>,
307}
308
309impl Serialize for KeyRecord {
310 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
311 where
312 S: Serializer,
313 {
314 let limits = KeyLimitsOut {
315 allowed_markets: self.allowed_markets.as_ref(),
316 allowed_symbols: self.allowed_symbols.as_ref(),
317 max_order_value: self.max_order_value,
318 max_daily_value: self.max_daily_value,
319 hours_window: self.hours_window.as_ref(),
320 max_orders_per_minute: self.max_orders_per_minute,
321 allowed_trd_sides: self.allowed_trd_sides.as_ref(),
322 allowed_acc_ids: self.raw_explicit_acc_ids.as_ref(),
323 allowed_card_nums: self.allowed_card_nums.as_ref(),
324 };
325 let limits = (!limits.is_empty()).then_some(limits);
326 KeyRecordOut {
327 id: &self.id,
328 hash: &self.hash,
329 scopes: &self.scopes,
330 limits,
331 expires_at: self.expires_at,
332 created_at: self.created_at,
333 note: self.note.as_ref(),
334 allowed_machines: self.allowed_machines.as_ref(),
335 }
336 .serialize(serializer)
337 }
338}
339
340impl KeyRecord {
341 #[must_use = "丢弃生成结果会丢失 plaintext; 调用方必须立即展示给用户"]
345 pub fn generate(
346 id: impl Into<String>,
347 scopes: HashSet<Scope>,
348 limits: Option<Limits>,
349 expires_at: Option<DateTime<Utc>>,
350 note: Option<String>,
351 ) -> (String, KeyRecord) {
352 Self::generate_with_machines(id, scopes, limits, expires_at, note, None)
353 }
354
355 #[must_use = "丢弃生成结果会丢失 plaintext; 调用方必须立即展示给用户"]
357 pub fn generate_with_machines(
358 id: impl Into<String>,
359 scopes: HashSet<Scope>,
360 limits: Option<Limits>,
361 expires_at: Option<DateTime<Utc>>,
362 note: Option<String>,
363 allowed_machines: Option<Vec<String>>,
364 ) -> (String, KeyRecord) {
365 let mut bytes = [0u8; 32];
366 rand::rng().fill_bytes(&mut bytes);
367 let plaintext = hex::encode(bytes);
368 let hash = format!(
369 "sha256:{}",
370 hex::encode(Sha256::digest(plaintext.as_bytes()))
371 );
372 let limits = limits.unwrap_or_default();
373 let raw_explicit_acc_ids = limits.allowed_acc_ids.clone();
375 let record = KeyRecord {
376 id: id.into(),
377 hash,
378 scopes,
379 allowed_markets: limits.allowed_markets,
380 allowed_symbols: limits.allowed_symbols,
381 max_order_value: limits.max_order_value,
382 max_daily_value: limits.max_daily_value,
383 hours_window: limits.hours_window,
384 max_orders_per_minute: limits.max_orders_per_minute,
385 allowed_trd_sides: limits.allowed_trd_sides,
386 allowed_acc_ids: limits.allowed_acc_ids,
387 allowed_card_nums: limits.allowed_card_nums,
388 expires_at,
389 created_at: Utc::now(),
390 note,
391 allowed_machines,
392 raw_explicit_acc_ids,
393 };
394 (plaintext, record)
395 }
396
397 pub fn check_machine(&self) -> Result<(), crate::machine::MachineError> {
399 crate::machine::check(&self.id, self.allowed_machines.as_deref())
400 }
401
402 #[must_use]
404 pub fn matches(&self, plaintext: &str) -> bool {
405 let computed = hash_plaintext(plaintext);
406 self.matches_hash(&computed)
407 }
408
409 pub(crate) fn matches_hash(&self, computed: &str) -> bool {
410 constant_time_eq_str(&self.hash, computed)
411 }
412
413 pub(crate) fn is_generated_plaintext_shape(plaintext: &str) -> bool {
417 plaintext.len() == 64
418 && plaintext
419 .bytes()
420 .all(|b| b.is_ascii_digit() || (b'a'..=b'f').contains(&b))
421 }
422
423 #[must_use]
425 pub fn is_expired(&self, now: DateTime<Utc>) -> bool {
426 self.expires_at.map(|t| now >= t).unwrap_or(false)
427 }
428
429 pub fn hours_range(&self) -> Result<Option<(NaiveTime, NaiveTime)>, String> {
431 let Some(s) = &self.hours_window else {
432 return Ok(None);
433 };
434 let (l, r) = s
435 .split_once('-')
436 .ok_or_else(|| format!("invalid hours_window {s:?}: expect HH:MM-HH:MM"))?;
437 let parse = |p: &str| {
438 NaiveTime::parse_from_str(p.trim(), "%H:%M")
439 .map_err(|e| format!("invalid time {p:?}: {e}"))
440 };
441 Ok(Some((parse(l)?, parse(r)?)))
442 }
443
444 #[must_use]
446 pub fn limits(&self) -> Limits {
447 Limits {
448 allowed_markets: self.allowed_markets.clone(),
449 allowed_symbols: self.allowed_symbols.clone(),
450 max_order_value: self.max_order_value,
451 max_daily_value: self.max_daily_value,
452 hours_window: self.hours_window.clone(),
453 max_orders_per_minute: self.max_orders_per_minute,
454 allowed_trd_sides: self.allowed_trd_sides.clone(),
455 allowed_acc_ids: self.allowed_acc_ids.clone(),
456 allowed_card_nums: self.allowed_card_nums.clone(),
457 }
458 }
459}
460
461impl crate::limits::LimitPolicy for KeyRecord {
462 fn allowed_markets(&self) -> Option<&HashSet<String>> {
463 self.allowed_markets.as_ref()
464 }
465
466 fn allowed_symbols(&self) -> Option<&HashSet<String>> {
467 self.allowed_symbols.as_ref()
468 }
469
470 fn max_order_value(&self) -> Option<f64> {
471 self.max_order_value
472 }
473
474 fn max_daily_value(&self) -> Option<f64> {
475 self.max_daily_value
476 }
477
478 fn hours_window(&self) -> Option<&str> {
479 self.hours_window.as_deref()
480 }
481
482 fn max_orders_per_minute(&self) -> Option<u32> {
483 self.max_orders_per_minute
484 }
485
486 fn allowed_trd_sides(&self) -> Option<&HashSet<String>> {
487 self.allowed_trd_sides.as_ref()
488 }
489
490 fn allowed_acc_ids(&self) -> Option<&HashSet<u64>> {
491 self.allowed_acc_ids.as_ref()
492 }
493
494 fn allowed_card_nums(&self) -> Option<&[String]> {
495 self.allowed_card_nums.as_deref()
496 }
497}
498
499fn constant_time_eq_str(a: &str, b: &str) -> bool {
500 let a = a.as_bytes();
501 let b = b.as_bytes();
502 if a.len() != b.len() {
503 return false;
504 }
505 let mut acc: u8 = 0;
506 for (x, y) in a.iter().zip(b.iter()) {
507 acc |= x ^ y;
508 }
509 acc == 0
510}
511
512#[must_use]
514pub fn hash_plaintext(plaintext: &str) -> String {
515 format!(
516 "sha256:{}",
517 hex::encode(Sha256::digest(plaintext.as_bytes()))
518 )
519}
520
521#[cfg(test)]
522mod tests;