1use std::collections::{HashMap, HashSet, hash_map::DefaultHasher};
2use std::hash::{Hash, Hasher};
3use std::num::NonZeroUsize;
4use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
5use std::sync::{Arc, OnceLock, RwLock};
6use std::time::{Duration, Instant};
7
8pub const DEFAULT_PROVIDER_NAME: &str = "default";
10
11use crate::error::{KnowReason, KnowledgeResult};
12use async_trait::async_trait;
13use lru::LruCache;
14use orion_error::conversion::ToStructError;
15use tokio::task;
16use wp_log::{debug_kdb, warn_kdb};
17use wp_model_core::model::{DataField, DataType, Value};
18
19use crate::loader::ProviderKind;
20use crate::mem::RowData;
21use crate::telemetry::{
22 CacheLayer, CacheOutcome, CacheTelemetryEvent, QueryTelemetryEvent, ReloadOutcome,
23 ReloadTelemetryEvent, telemetry, telemetry_enabled,
24};
25
26#[derive(Debug, Clone, PartialEq, Eq, Hash)]
27pub struct DatasourceId(pub String);
28
29impl DatasourceId {
30 pub fn from_seed(kind: ProviderKind, seed: &str) -> Self {
31 let mut hasher = DefaultHasher::new();
32 seed.hash(&mut hasher);
33 let kind_str = match kind {
34 ProviderKind::SqliteAuthority => "sqlite",
35 ProviderKind::Postgres => "postgres",
36 ProviderKind::Mysql => "mysql",
37 ProviderKind::Redis => "redis",
38 };
39 Self(format!("{kind_str}:{:016x}", hasher.finish()))
40 }
41}
42
43#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
44pub struct Generation(pub u64);
45
46#[derive(Debug, Clone)]
47pub enum QueryMode {
48 Many,
49 FirstRow,
50}
51
52#[derive(Debug, Clone, Copy)]
53pub enum CachePolicy {
54 Bypass,
55 UseGlobal,
56 UseCallScope,
57}
58
59#[derive(Debug, Clone)]
60pub enum QueryValue {
61 Null,
62 Bool(bool),
63 Int(i64),
64 Float(f64),
65 Text(String),
66}
67
68#[derive(Debug, Clone)]
69pub struct QueryParam {
70 pub name: String,
71 pub value: QueryValue,
72}
73
74#[derive(Debug, Clone)]
75pub struct QueryRequest {
76 pub sql: String,
77 pub params: Vec<QueryParam>,
78 pub mode: QueryMode,
79 pub cache_policy: CachePolicy,
80}
81
82impl QueryRequest {
83 pub fn many(
84 sql: impl Into<String>,
85 params: Vec<QueryParam>,
86 cache_policy: CachePolicy,
87 ) -> Self {
88 Self {
89 sql: sql.into(),
90 params,
91 mode: QueryMode::Many,
92 cache_policy,
93 }
94 }
95
96 pub fn first_row(
97 sql: impl Into<String>,
98 params: Vec<QueryParam>,
99 cache_policy: CachePolicy,
100 ) -> Self {
101 Self {
102 sql: sql.into(),
103 params,
104 mode: QueryMode::FirstRow,
105 cache_policy,
106 }
107 }
108}
109
110#[derive(Debug, Clone)]
111pub enum QueryResponse {
112 Rows(Vec<RowData>),
113 Row(RowData),
114}
115
116impl QueryResponse {
117 pub fn into_rows(self) -> Vec<RowData> {
118 match self {
119 QueryResponse::Rows(rows) => rows,
120 QueryResponse::Row(row) => vec![row],
121 }
122 }
123
124 pub fn into_row(self) -> RowData {
125 match self {
126 QueryResponse::Rows(rows) => rows.into_iter().next().unwrap_or_default(),
127 QueryResponse::Row(row) => row,
128 }
129 }
130}
131
132#[async_trait]
133pub trait ProviderExecutor: Send + Sync {
134 fn query(&self, sql: &str) -> KnowledgeResult<Vec<RowData>>;
135 fn query_fields(&self, sql: &str, params: &[DataField]) -> KnowledgeResult<Vec<RowData>>;
136 fn query_row(&self, sql: &str) -> KnowledgeResult<RowData>;
137 fn query_named_fields(&self, sql: &str, params: &[DataField]) -> KnowledgeResult<RowData>;
138
139 async fn query_async(&self, sql: &str) -> KnowledgeResult<Vec<RowData>> {
140 self.query(sql)
141 }
142
143 async fn query_fields_async(
144 &self,
145 sql: &str,
146 params: &[DataField],
147 ) -> KnowledgeResult<Vec<RowData>> {
148 self.query_fields(sql, params)
149 }
150
151 async fn query_row_async(&self, sql: &str) -> KnowledgeResult<RowData> {
152 self.query_row(sql)
153 }
154
155 async fn query_named_fields_async(
156 &self,
157 sql: &str,
158 params: &[DataField],
159 ) -> KnowledgeResult<RowData> {
160 self.query_named_fields(sql, params)
161 }
162}
163
164#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
165pub enum QueryModeTag {
166 Many,
167 FirstRow,
168}
169
170#[derive(Debug, Clone, PartialEq, Eq, Hash)]
171pub struct ResultCacheKey {
172 pub datasource_id: DatasourceId,
173 pub generation: Generation,
174 pub query_hash: u64,
175 pub params_hash: u64,
176 pub mode: QueryModeTag,
177}
178
179pub struct ProviderHandle {
180 pub provider: Arc<dyn ProviderExecutor>,
181 pub datasource_id: DatasourceId,
182 pub generation: Generation,
183 pub kind: ProviderKind,
184}
185
186#[derive(Debug, Clone)]
187pub struct RuntimeSnapshot {
188 pub provider_kind: Option<ProviderKind>,
189 pub datasource_id: Option<DatasourceId>,
190 pub generation: Option<Generation>,
191 pub result_cache_enabled: bool,
192 pub result_cache_len: usize,
193 pub result_cache_capacity: usize,
194 pub result_cache_ttl_ms: u64,
195 pub metadata_cache_len: usize,
196 pub metadata_cache_capacity: usize,
197 pub result_cache_hits: u64,
198 pub result_cache_misses: u64,
199 pub metadata_cache_hits: u64,
200 pub metadata_cache_misses: u64,
201 pub local_cache_hits: u64,
202 pub local_cache_misses: u64,
203 pub reload_successes: u64,
204 pub reload_failures: u64,
205}
206
207#[derive(Debug, Clone)]
208pub struct MetadataCacheScope {
209 pub datasource_id: DatasourceId,
210 pub generation: Generation,
211}
212
213#[derive(Debug, Clone, Copy)]
214pub struct ResultCacheConfig {
215 pub enabled: bool,
216 pub capacity: usize,
217 pub ttl: Duration,
218}
219
220impl Default for ResultCacheConfig {
221 fn default() -> Self {
222 Self {
223 enabled: true,
224 capacity: 1024,
225 ttl: Duration::from_millis(30_000),
226 }
227 }
228}
229
230#[derive(Debug, Clone)]
231struct CachedQueryResponse {
232 response: Arc<QueryResponse>,
233 cached_at: Instant,
234}
235
236#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
241pub(crate) enum RedisCmdTag {
242 BfExists,
243 HGet,
244 Get,
245 SetExists,
246}
247
248#[derive(Debug, Clone, PartialEq, Eq, Hash)]
249pub(crate) struct RedisCacheKey {
250 pub generation: u64,
251 pub cmd_tag: RedisCmdTag,
252 pub key_hash: u64,
253 pub args_hash: u64,
254}
255
256#[derive(Debug, Clone)]
257pub(crate) enum CachedRedisValue {
258 Bool(bool),
259 OptString(Option<String>),
260}
261
262#[derive(Debug, Clone)]
263struct CachedRedisEntry {
264 value: CachedRedisValue,
265 cached_at: Instant,
266 ttl_ms: u64, }
268
269pub struct KnowledgeRuntime {
270 provider: RwLock<Option<Arc<ProviderHandle>>>,
271 named_providers: RwLock<HashMap<String, Arc<ProviderHandle>>>,
273 named_provider_epoch: AtomicU64,
275 next_generation: AtomicU64,
276 provider_epoch: AtomicU64,
277 current_generation_value: AtomicU64,
278 result_cache_config: RwLock<ResultCacheConfig>,
279 result_cache_enabled: AtomicBool,
280 result_cache_ttl_ms: AtomicU64,
281 result_cache: RwLock<LruCache<ResultCacheKey, CachedQueryResponse>>,
282 result_cache_hits: AtomicU64,
283 result_cache_misses: AtomicU64,
284 metadata_cache_hits: AtomicU64,
285 metadata_cache_misses: AtomicU64,
286 local_cache_hits: AtomicU64,
287 local_cache_misses: AtomicU64,
288 reload_successes: AtomicU64,
289 reload_failures: AtomicU64,
290 redis_cache: RwLock<LruCache<RedisCacheKey, CachedRedisEntry>>,
291 redis_cache_hits: AtomicU64,
292 redis_cache_misses: AtomicU64,
293 redis_global_enabled: AtomicBool,
294}
295
296impl KnowledgeRuntime {
297 pub fn new(result_cache_capacity: usize) -> Self {
298 let config = ResultCacheConfig {
299 capacity: result_cache_capacity.max(1),
300 ..ResultCacheConfig::default()
301 };
302 let capacity = NonZeroUsize::new(config.capacity).expect("non-zero capacity");
303 Self {
304 provider: RwLock::new(None),
305 named_providers: RwLock::new(HashMap::new()),
306 named_provider_epoch: AtomicU64::new(0),
307 next_generation: AtomicU64::new(0),
308 provider_epoch: AtomicU64::new(0),
309 current_generation_value: AtomicU64::new(0),
310 result_cache_config: RwLock::new(config),
311 result_cache_enabled: AtomicBool::new(config.enabled),
312 result_cache_ttl_ms: AtomicU64::new(config.ttl.as_millis() as u64),
313 result_cache: RwLock::new(LruCache::new(capacity)),
314 result_cache_hits: AtomicU64::new(0),
315 result_cache_misses: AtomicU64::new(0),
316 metadata_cache_hits: AtomicU64::new(0),
317 metadata_cache_misses: AtomicU64::new(0),
318 local_cache_hits: AtomicU64::new(0),
319 local_cache_misses: AtomicU64::new(0),
320 reload_successes: AtomicU64::new(0),
321 reload_failures: AtomicU64::new(0),
322 redis_cache: RwLock::new(LruCache::new(capacity)),
323 redis_cache_hits: AtomicU64::new(0),
324 redis_cache_misses: AtomicU64::new(0),
325 redis_global_enabled: AtomicBool::new(true),
326 }
327 }
328
329 pub fn install_provider<F>(
330 &self,
331 kind: ProviderKind,
332 datasource_id: DatasourceId,
333 build: F,
334 ) -> KnowledgeResult<Generation>
335 where
336 F: FnOnce(Generation) -> KnowledgeResult<Arc<dyn ProviderExecutor>>,
337 {
338 self.install_provider_named(DEFAULT_PROVIDER_NAME, kind, datasource_id, build, true)
339 }
340
341 pub fn install_provider_named<F>(
345 &self,
346 name: &str,
347 kind: ProviderKind,
348 datasource_id: DatasourceId,
349 build: F,
350 is_default: bool,
351 ) -> KnowledgeResult<Generation>
352 where
353 F: FnOnce(Generation) -> KnowledgeResult<Arc<dyn ProviderExecutor>>,
354 {
355 let generation = Generation(self.next_generation.fetch_add(1, Ordering::SeqCst) + 1);
356 let previous = self
357 .provider
358 .read()
359 .ok()
360 .and_then(|guard| guard.as_ref().cloned());
361 debug_kdb!(
362 "[kdb] reload provider start name={} kind={kind:?} datasource_id={} target_generation={} previous_generation={}",
363 name,
364 datasource_id.0,
365 generation.0,
366 previous
367 .as_ref()
368 .map(|handle| handle.generation.0.to_string())
369 .unwrap_or_else(|| "none".to_string())
370 );
371 let provider = match build(generation) {
372 Ok(provider) => provider,
373 Err(err) => {
374 self.reload_failures.fetch_add(1, Ordering::Relaxed);
375 warn_kdb!(
376 "[kdb] reload provider failed name={} kind={kind:?} datasource_id={} target_generation={} err={}",
377 name,
378 datasource_id.0,
379 generation.0,
380 err
381 );
382 if telemetry_enabled() {
383 telemetry().on_reload(&ReloadTelemetryEvent {
384 outcome: ReloadOutcome::Failure,
385 provider_kind: kind.clone(),
386 });
387 }
388 return Err(err);
389 }
390 };
391 debug_kdb!(
392 "[kdb] install provider name={} kind={kind:?} datasource_id={} generation={}",
393 name,
394 datasource_id.0,
395 generation.0
396 );
397 let kind_for_handle = kind.clone();
398 let datasource_id_for_handle = datasource_id.clone();
399 let handle = Arc::new(ProviderHandle {
400 provider,
401 datasource_id: datasource_id_for_handle,
402 generation,
403 kind: kind_for_handle,
404 });
405 {
406 let mut guard = self
407 .named_providers
408 .write()
409 .expect("runtime named provider lock poisoned");
410 guard.insert(name.to_string(), handle.clone());
411 }
412 self.named_provider_epoch.fetch_add(1, Ordering::AcqRel);
413 if is_default {
414 self.provider_epoch.fetch_add(1, Ordering::AcqRel);
415 {
416 let mut guard = self
417 .provider
418 .write()
419 .expect("runtime provider lock poisoned");
420 *guard = Some(handle);
421 }
422 self.current_generation_value
423 .store(generation.0, Ordering::Release);
424 self.provider_epoch.fetch_add(1, Ordering::Release);
425 }
426 self.reload_successes.fetch_add(1, Ordering::Relaxed);
427 if telemetry_enabled() {
428 telemetry().on_reload(&ReloadTelemetryEvent {
429 outcome: ReloadOutcome::Success,
430 provider_kind: kind.clone(),
431 });
432 }
433 debug_kdb!(
434 "[kdb] reload provider success name={} kind={kind:?} datasource_id={} generation={}",
435 name,
436 datasource_id.0,
437 generation.0
438 );
439 Ok(generation)
440 }
441
442 pub fn configure_result_cache(&self, enabled: bool, capacity: usize, ttl: Duration) {
443 let new_config = ResultCacheConfig {
444 enabled,
445 capacity: capacity.max(1),
446 ttl: ttl.max(Duration::from_millis(1)),
447 };
448 let mut should_reset_cache = false;
449 {
450 let mut guard = self
451 .result_cache_config
452 .write()
453 .expect("runtime result cache config lock poisoned");
454 if guard.capacity != new_config.capacity || (!new_config.enabled && guard.enabled) {
455 should_reset_cache = true;
456 }
457 *guard = new_config;
458 }
459 self.result_cache_enabled
460 .store(new_config.enabled, Ordering::Relaxed);
461 self.result_cache_ttl_ms.store(
462 new_config.ttl.as_millis().min(u128::from(u64::MAX)) as u64,
463 Ordering::Relaxed,
464 );
465
466 if should_reset_cache {
467 let mut cache = self
468 .result_cache
469 .write()
470 .expect("runtime result cache lock poisoned");
471 *cache = LruCache::new(
472 NonZeroUsize::new(new_config.capacity).expect("non-zero result cache capacity"),
473 );
474 }
475 }
476
477 pub fn configure_redis_cache(&self, global_enabled: bool, capacity: usize) {
478 let new_capacity =
479 NonZeroUsize::new(capacity.max(1)).expect("non-zero redis cache capacity");
480 if let Ok(mut cache) = self.redis_cache.write() {
481 *cache = LruCache::new(new_capacity);
482 }
483 self.redis_global_enabled
484 .store(global_enabled, Ordering::Relaxed);
485 }
486
487 pub fn current_generation(&self) -> Option<Generation> {
488 let epoch_before = self.provider_epoch.load(Ordering::Acquire);
489 if epoch_before % 2 == 1 {
490 return self.current_generation_from_provider();
491 }
492 let generation = self.current_generation_value.load(Ordering::Acquire);
493 let epoch_after = self.provider_epoch.load(Ordering::Acquire);
494 if epoch_before != epoch_after {
495 return self.current_generation_from_provider();
496 }
497 match generation {
498 0 => None,
499 generation => Some(Generation(generation)),
500 }
501 }
502
503 pub fn snapshot(&self) -> RuntimeSnapshot {
504 let provider = self
505 .provider
506 .read()
507 .ok()
508 .and_then(|guard| guard.as_ref().cloned());
509 let result_cache_config = self
510 .result_cache_config
511 .read()
512 .map(|guard| *guard)
513 .unwrap_or_default();
514 let (result_cache_len, result_cache_capacity) = self
515 .result_cache
516 .read()
517 .map(|cache| (cache.len(), cache.cap().get()))
518 .unwrap_or((0, 0));
519 let (metadata_cache_len, metadata_cache_capacity) =
520 crate::mem::query_util::column_metadata_cache_snapshot();
521 RuntimeSnapshot {
522 provider_kind: provider.as_ref().map(|handle| handle.kind.clone()),
523 datasource_id: provider.as_ref().map(|handle| handle.datasource_id.clone()),
524 generation: provider.as_ref().map(|handle| handle.generation),
525 result_cache_enabled: result_cache_config.enabled,
526 result_cache_len,
527 result_cache_capacity,
528 result_cache_ttl_ms: result_cache_config.ttl.as_millis() as u64,
529 metadata_cache_len,
530 metadata_cache_capacity,
531 result_cache_hits: self.result_cache_hits.load(Ordering::Relaxed),
532 result_cache_misses: self.result_cache_misses.load(Ordering::Relaxed),
533 metadata_cache_hits: self.metadata_cache_hits.load(Ordering::Relaxed),
534 metadata_cache_misses: self.metadata_cache_misses.load(Ordering::Relaxed),
535 local_cache_hits: self.local_cache_hits.load(Ordering::Relaxed),
536 local_cache_misses: self.local_cache_misses.load(Ordering::Relaxed),
537 reload_successes: self.reload_successes.load(Ordering::Relaxed),
538 reload_failures: self.reload_failures.load(Ordering::Relaxed),
539 }
540 }
541
542 pub fn current_metadata_scope(&self) -> MetadataCacheScope {
543 self.provider
544 .read()
545 .ok()
546 .and_then(|guard| guard.as_ref().cloned())
547 .map(|handle| MetadataCacheScope {
548 datasource_id: handle.datasource_id.clone(),
549 generation: handle.generation,
550 })
551 .unwrap_or_else(|| MetadataCacheScope {
552 datasource_id: DatasourceId("sqlite:standalone".to_string()),
553 generation: Generation(0),
554 })
555 }
556
557 pub fn current_provider_kind(&self) -> Option<ProviderKind> {
558 self.provider
559 .read()
560 .ok()
561 .and_then(|guard| guard.as_ref().map(|handle| handle.kind.clone()))
562 }
563
564 pub fn provider_by_name(&self, name: &str) -> KnowledgeResult<Arc<ProviderHandle>> {
566 self.named_providers
567 .read()
568 .expect("runtime named provider lock poisoned")
569 .get(name)
570 .cloned()
571 .ok_or_else(|| {
572 KnowReason::from_logic()
573 .to_err()
574 .with_detail(format!("knowledge provider '{name}' not initialized"))
575 })
576 }
577
578 pub fn provider_exists(&self, name: &str) -> bool {
579 self.named_providers
580 .read()
581 .ok()
582 .is_some_and(|guard| guard.contains_key(name))
583 }
584
585 pub fn provider_names(&self) -> Vec<String> {
586 self.named_providers
587 .read()
588 .ok()
589 .map(|guard| guard.keys().cloned().collect())
590 .unwrap_or_default()
591 }
592
593 pub fn prune_named_providers_except(&self, keep: &HashSet<String>) {
596 if let Ok(mut guard) = self.named_providers.write() {
597 let before = guard.len();
598 guard.retain(|name, _| keep.contains(name));
599 if guard.len() != before {
600 self.named_provider_epoch.fetch_add(1, Ordering::AcqRel);
601 }
602 }
603 }
604
605 pub fn named_provider_epoch(&self) -> u64 {
607 self.named_provider_epoch.load(Ordering::Acquire)
608 }
609
610 pub fn record_result_cache_hit(&self) {
611 self.result_cache_hits.fetch_add(1, Ordering::Relaxed);
612 }
613
614 pub fn record_result_cache_miss(&self) {
615 self.result_cache_misses.fetch_add(1, Ordering::Relaxed);
616 }
617
618 pub fn record_metadata_cache_hit(&self) {
619 self.metadata_cache_hits.fetch_add(1, Ordering::Relaxed);
620 }
621
622 pub fn record_metadata_cache_miss(&self) {
623 self.metadata_cache_misses.fetch_add(1, Ordering::Relaxed);
624 }
625
626 pub fn record_local_cache_hit(&self) {
627 self.local_cache_hits.fetch_add(1, Ordering::Relaxed);
628 }
629
630 pub fn record_local_cache_miss(&self) {
631 self.local_cache_misses.fetch_add(1, Ordering::Relaxed);
632 }
633
634 pub fn execute(&self, req: &QueryRequest) -> KnowledgeResult<QueryResponse> {
635 let handle = self.current_handle()?;
636 self.execute_with_handle(&handle, req)
637 }
638
639 pub fn execute_for(&self, name: &str, req: &QueryRequest) -> KnowledgeResult<QueryResponse> {
641 let handle = self.provider_by_name(name)?;
642 self.execute_with_handle(&handle, req)
643 }
644
645 fn execute_with_handle(
646 &self,
647 handle: &Arc<ProviderHandle>,
648 req: &QueryRequest,
649 ) -> KnowledgeResult<QueryResponse> {
650 let use_global_cache =
651 matches!(req.cache_policy, CachePolicy::UseGlobal) && self.result_cache_enabled();
652 if use_global_cache && let Some(hit) = self.fetch_result_cache(handle, req) {
653 self.record_result_cache_hit();
654 if telemetry_enabled() {
655 telemetry().on_cache(&CacheTelemetryEvent {
656 layer: CacheLayer::Result,
657 outcome: CacheOutcome::Hit,
658 provider_kind: Some(handle.kind.clone()),
659 });
660 }
661 debug_kdb!(
662 "[kdb] global result cache hit kind={:?} generation={}",
663 handle.kind,
664 handle.generation.0
665 );
666 return Ok(hit);
667 }
668 if use_global_cache {
669 self.record_result_cache_miss();
670 if telemetry_enabled() {
671 telemetry().on_cache(&CacheTelemetryEvent {
672 layer: CacheLayer::Result,
673 outcome: CacheOutcome::Miss,
674 provider_kind: Some(handle.kind.clone()),
675 });
676 }
677 debug_kdb!(
678 "[kdb] global result cache miss kind={:?} generation={}",
679 handle.kind,
680 handle.generation.0
681 );
682 }
683
684 let params = params_to_fields(&req.params);
685 let mode_tag = query_mode_tag(&req.mode);
686 let started = Instant::now();
687 debug_kdb!(
688 "[kdb] execute query kind={:?} generation={} mode={:?} cache_policy={:?}",
689 handle.kind,
690 handle.generation.0,
691 req.mode,
692 req.cache_policy
693 );
694 let response = match match req.mode {
695 QueryMode::Many => {
696 if params.is_empty() {
697 handle.provider.query(&req.sql).map(QueryResponse::Rows)
698 } else {
699 handle
700 .provider
701 .query_fields(&req.sql, ¶ms)
702 .map(QueryResponse::Rows)
703 }
704 }
705 QueryMode::FirstRow => {
706 if params.is_empty() {
707 handle.provider.query_row(&req.sql).map(QueryResponse::Row)
708 } else {
709 handle
710 .provider
711 .query_named_fields(&req.sql, ¶ms)
712 .map(QueryResponse::Row)
713 }
714 }
715 } {
716 Ok(response) => {
717 if telemetry_enabled() {
718 telemetry().on_query(&QueryTelemetryEvent {
719 provider_kind: handle.kind.clone(),
720 mode: mode_tag,
721 success: true,
722 elapsed: started.elapsed(),
723 });
724 }
725 response
726 }
727 Err(err) => {
728 if telemetry_enabled() {
729 telemetry().on_query(&QueryTelemetryEvent {
730 provider_kind: handle.kind.clone(),
731 mode: mode_tag,
732 success: false,
733 elapsed: started.elapsed(),
734 });
735 }
736 return Err(err);
737 }
738 };
739
740 if use_global_cache {
741 self.save_result_cache(handle, req, response.clone());
742 debug_kdb!(
743 "[kdb] global result cache store kind={:?} generation={}",
744 handle.kind,
745 handle.generation.0
746 );
747 }
748
749 Ok(response)
750 }
751
752 pub fn execute_first_row_fields(
753 &self,
754 sql: &str,
755 params: &[DataField],
756 cache_policy: CachePolicy,
757 ) -> KnowledgeResult<RowData> {
758 let handle = self.current_handle()?;
759 self.execute_first_row_fields_with_handle(&handle, sql, params, cache_policy)
760 }
761
762 pub fn execute_first_row_fields_for(
764 &self,
765 name: &str,
766 sql: &str,
767 params: &[DataField],
768 cache_policy: CachePolicy,
769 ) -> KnowledgeResult<RowData> {
770 let handle = self.provider_by_name(name)?;
771 self.execute_first_row_fields_with_handle(&handle, sql, params, cache_policy)
772 }
773
774 fn execute_first_row_fields_with_handle(
775 &self,
776 handle: &Arc<ProviderHandle>,
777 sql: &str,
778 params: &[DataField],
779 cache_policy: CachePolicy,
780 ) -> KnowledgeResult<RowData> {
781 let use_global_cache =
782 matches!(cache_policy, CachePolicy::UseGlobal) && self.result_cache_enabled();
783 if use_global_cache
784 && let Some(hit) = self.fetch_result_cache_by_key(result_cache_key_fields(
785 handle,
786 sql,
787 params,
788 QueryModeTag::FirstRow,
789 ))
790 {
791 self.record_result_cache_hit();
792 if telemetry_enabled() {
793 telemetry().on_cache(&CacheTelemetryEvent {
794 layer: CacheLayer::Result,
795 outcome: CacheOutcome::Hit,
796 provider_kind: Some(handle.kind.clone()),
797 });
798 }
799 return Ok(hit.into_row());
800 }
801 if use_global_cache {
802 self.record_result_cache_miss();
803 if telemetry_enabled() {
804 telemetry().on_cache(&CacheTelemetryEvent {
805 layer: CacheLayer::Result,
806 outcome: CacheOutcome::Miss,
807 provider_kind: Some(handle.kind.clone()),
808 });
809 }
810 }
811
812 let started = Instant::now();
813 let row = if params.is_empty() {
814 handle.provider.query_row(sql)
815 } else {
816 handle.provider.query_named_fields(sql, params)
817 };
818 let row = match row {
819 Ok(row) => {
820 if telemetry_enabled() {
821 telemetry().on_query(&QueryTelemetryEvent {
822 provider_kind: handle.kind.clone(),
823 mode: QueryModeTag::FirstRow,
824 success: true,
825 elapsed: started.elapsed(),
826 });
827 }
828 row
829 }
830 Err(err) => {
831 if telemetry_enabled() {
832 telemetry().on_query(&QueryTelemetryEvent {
833 provider_kind: handle.kind.clone(),
834 mode: QueryModeTag::FirstRow,
835 success: false,
836 elapsed: started.elapsed(),
837 });
838 }
839 return Err(err);
840 }
841 };
842
843 if use_global_cache {
844 self.save_result_cache_by_key(
845 result_cache_key_fields(handle, sql, params, QueryModeTag::FirstRow),
846 QueryResponse::Row(row.clone()),
847 );
848 }
849
850 Ok(row)
851 }
852
853 pub async fn execute_async(&self, req: &QueryRequest) -> KnowledgeResult<QueryResponse> {
854 let handle = self.current_handle()?;
855 if matches!(handle.kind, ProviderKind::SqliteAuthority) {
856 let handle = handle.clone();
857 let req = req.clone();
858 return task::spawn_blocking(move || runtime().execute_with_handle(&handle, &req))
859 .await
860 .map_err(|err| {
861 KnowReason::from_logic()
862 .to_err()
863 .with_detail(format!("knowledge async sqlite query join failed: {err}"))
864 })?;
865 }
866 self.execute_async_with_handle(&handle, req).await
867 }
868
869 pub async fn execute_for_async(
870 &self,
871 name: &str,
872 req: &QueryRequest,
873 ) -> KnowledgeResult<QueryResponse> {
874 let handle = self.provider_by_name(name)?;
875 self.execute_async_with_handle(&handle, req).await
876 }
877
878 async fn execute_async_with_handle(
879 &self,
880 handle: &Arc<ProviderHandle>,
881 req: &QueryRequest,
882 ) -> KnowledgeResult<QueryResponse> {
883 let use_global_cache =
884 matches!(req.cache_policy, CachePolicy::UseGlobal) && self.result_cache_enabled();
885 if use_global_cache && let Some(hit) = self.fetch_result_cache(handle, req) {
886 self.record_result_cache_hit();
887 if telemetry_enabled() {
888 telemetry().on_cache(&CacheTelemetryEvent {
889 layer: CacheLayer::Result,
890 outcome: CacheOutcome::Hit,
891 provider_kind: Some(handle.kind.clone()),
892 });
893 }
894 return Ok(hit);
895 }
896 if use_global_cache {
897 self.record_result_cache_miss();
898 if telemetry_enabled() {
899 telemetry().on_cache(&CacheTelemetryEvent {
900 layer: CacheLayer::Result,
901 outcome: CacheOutcome::Miss,
902 provider_kind: Some(handle.kind.clone()),
903 });
904 }
905 }
906
907 let params = params_to_fields(&req.params);
908 let mode_tag = query_mode_tag(&req.mode);
909 let started = Instant::now();
910 let response = match req.mode {
911 QueryMode::Many => {
912 if params.is_empty() {
913 handle
914 .provider
915 .query_async(&req.sql)
916 .await
917 .map(QueryResponse::Rows)
918 } else {
919 handle
920 .provider
921 .query_fields_async(&req.sql, ¶ms)
922 .await
923 .map(QueryResponse::Rows)
924 }
925 }
926 QueryMode::FirstRow => {
927 if params.is_empty() {
928 handle
929 .provider
930 .query_row_async(&req.sql)
931 .await
932 .map(QueryResponse::Row)
933 } else {
934 handle
935 .provider
936 .query_named_fields_async(&req.sql, ¶ms)
937 .await
938 .map(QueryResponse::Row)
939 }
940 }
941 };
942 let response = match response {
943 Ok(response) => {
944 if telemetry_enabled() {
945 telemetry().on_query(&QueryTelemetryEvent {
946 provider_kind: handle.kind.clone(),
947 mode: mode_tag,
948 success: true,
949 elapsed: started.elapsed(),
950 });
951 }
952 response
953 }
954 Err(err) => {
955 if telemetry_enabled() {
956 telemetry().on_query(&QueryTelemetryEvent {
957 provider_kind: handle.kind.clone(),
958 mode: mode_tag,
959 success: false,
960 elapsed: started.elapsed(),
961 });
962 }
963 return Err(err);
964 }
965 };
966
967 if use_global_cache {
968 self.save_result_cache(handle, req, response.clone());
969 }
970
971 Ok(response)
972 }
973
974 pub async fn execute_first_row_fields_async(
975 &self,
976 sql: &str,
977 params: &[DataField],
978 cache_policy: CachePolicy,
979 ) -> KnowledgeResult<RowData> {
980 let handle = self.current_handle()?;
981 if matches!(handle.kind, ProviderKind::SqliteAuthority) {
982 let handle = handle.clone();
983 let sql = sql.to_string();
984 let params = params.to_vec();
985 return task::spawn_blocking(move || {
986 runtime().execute_first_row_fields_with_handle(&handle, &sql, ¶ms, cache_policy)
987 })
988 .await
989 .map_err(|err| {
990 KnowReason::from_logic().to_err().with_detail(format!(
991 "knowledge async sqlite first-row query join failed: {err}"
992 ))
993 })?;
994 }
995 self.execute_first_row_fields_async_with_handle(&handle, sql, params, cache_policy)
996 .await
997 }
998
999 pub async fn execute_first_row_fields_for_async(
1000 &self,
1001 name: &str,
1002 sql: &str,
1003 params: &[DataField],
1004 cache_policy: CachePolicy,
1005 ) -> KnowledgeResult<RowData> {
1006 let handle = self.provider_by_name(name)?;
1007 self.execute_first_row_fields_async_with_handle(&handle, sql, params, cache_policy)
1008 .await
1009 }
1010
1011 async fn execute_first_row_fields_async_with_handle(
1012 &self,
1013 handle: &Arc<ProviderHandle>,
1014 sql: &str,
1015 params: &[DataField],
1016 cache_policy: CachePolicy,
1017 ) -> KnowledgeResult<RowData> {
1018 let use_global_cache =
1019 matches!(cache_policy, CachePolicy::UseGlobal) && self.result_cache_enabled();
1020 if use_global_cache
1021 && let Some(hit) = self.fetch_result_cache_by_key(result_cache_key_fields(
1022 handle,
1023 sql,
1024 params,
1025 QueryModeTag::FirstRow,
1026 ))
1027 {
1028 self.record_result_cache_hit();
1029 if telemetry_enabled() {
1030 telemetry().on_cache(&CacheTelemetryEvent {
1031 layer: CacheLayer::Result,
1032 outcome: CacheOutcome::Hit,
1033 provider_kind: Some(handle.kind.clone()),
1034 });
1035 }
1036 return Ok(hit.into_row());
1037 }
1038 if use_global_cache {
1039 self.record_result_cache_miss();
1040 if telemetry_enabled() {
1041 telemetry().on_cache(&CacheTelemetryEvent {
1042 layer: CacheLayer::Result,
1043 outcome: CacheOutcome::Miss,
1044 provider_kind: Some(handle.kind.clone()),
1045 });
1046 }
1047 }
1048
1049 let started = Instant::now();
1050 let row = if params.is_empty() {
1051 handle.provider.query_row_async(sql).await
1052 } else {
1053 handle.provider.query_named_fields_async(sql, params).await
1054 };
1055 let row = match row {
1056 Ok(row) => {
1057 if telemetry_enabled() {
1058 telemetry().on_query(&QueryTelemetryEvent {
1059 provider_kind: handle.kind.clone(),
1060 mode: QueryModeTag::FirstRow,
1061 success: true,
1062 elapsed: started.elapsed(),
1063 });
1064 }
1065 row
1066 }
1067 Err(err) => {
1068 if telemetry_enabled() {
1069 telemetry().on_query(&QueryTelemetryEvent {
1070 provider_kind: handle.kind.clone(),
1071 mode: QueryModeTag::FirstRow,
1072 success: false,
1073 elapsed: started.elapsed(),
1074 });
1075 }
1076 return Err(err);
1077 }
1078 };
1079
1080 if use_global_cache {
1081 self.save_result_cache_by_key(
1082 result_cache_key_fields(handle, sql, params, QueryModeTag::FirstRow),
1083 QueryResponse::Row(row.clone()),
1084 );
1085 }
1086
1087 Ok(row)
1088 }
1089
1090 fn current_handle(&self) -> KnowledgeResult<Arc<ProviderHandle>> {
1091 self.provider
1092 .read()
1093 .expect("runtime provider lock poisoned")
1094 .clone()
1095 .ok_or_else(|| {
1096 KnowReason::from_logic()
1097 .to_err()
1098 .with_detail("knowledge provider not initialized")
1099 })
1100 }
1101
1102 fn current_generation_from_provider(&self) -> Option<Generation> {
1103 self.provider
1104 .read()
1105 .ok()
1106 .and_then(|guard| guard.as_ref().map(|handle| handle.generation))
1107 }
1108
1109 fn fetch_result_cache(
1110 &self,
1111 handle: &ProviderHandle,
1112 req: &QueryRequest,
1113 ) -> Option<QueryResponse> {
1114 self.fetch_result_cache_by_key(result_cache_key(handle, req))
1115 }
1116
1117 fn fetch_result_cache_by_key(&self, key: ResultCacheKey) -> Option<QueryResponse> {
1118 if !self.result_cache_enabled() {
1119 return None;
1120 }
1121 let cached = self
1122 .result_cache
1123 .read()
1124 .ok()
1125 .and_then(|cache| cache.peek(&key).cloned())?;
1126 if cached.cached_at.elapsed() > self.result_cache_ttl() {
1127 if let Ok(mut cache) = self.result_cache.write() {
1128 let _ = cache.pop(&key);
1129 }
1130 return None;
1131 }
1132 Some((*cached.response).clone())
1133 }
1134
1135 fn save_result_cache(
1136 &self,
1137 handle: &ProviderHandle,
1138 req: &QueryRequest,
1139 response: QueryResponse,
1140 ) {
1141 self.save_result_cache_by_key(result_cache_key(handle, req), response);
1142 }
1143
1144 fn save_result_cache_by_key(&self, key: ResultCacheKey, response: QueryResponse) {
1145 if let Ok(mut cache) = self.result_cache.write() {
1146 cache.put(
1147 key,
1148 CachedQueryResponse {
1149 response: Arc::new(response),
1150 cached_at: Instant::now(),
1151 },
1152 );
1153 }
1154 }
1155
1156 fn redis_cache_enabled(&self) -> bool {
1161 self.result_cache_enabled.load(Ordering::Acquire)
1162 }
1163
1164 #[allow(dead_code)]
1165 fn redis_cache_ttl(&self) -> Duration {
1166 Duration::from_millis(self.result_cache_ttl_ms.load(Ordering::Acquire))
1167 }
1168
1169 fn fetch_redis_cache(&self, key: &RedisCacheKey) -> Option<CachedRedisEntry> {
1170 if !self.redis_cache_enabled() {
1171 return None;
1172 }
1173 let entry = self
1174 .redis_cache
1175 .read()
1176 .ok()
1177 .and_then(|cache| cache.peek(key).cloned())?;
1178 if entry.ttl_ms > 0 && entry.cached_at.elapsed() > Duration::from_millis(entry.ttl_ms) {
1180 if let Ok(mut cache) = self.redis_cache.write() {
1181 let _ = cache.pop(key);
1182 }
1183 return None;
1184 }
1185 Some(entry)
1186 }
1187
1188 fn save_redis_cache(&self, key: RedisCacheKey, entry: CachedRedisEntry) {
1189 if !self.redis_cache_enabled() {
1190 return;
1191 }
1192 if let Ok(mut cache) = self.redis_cache.write() {
1193 cache.put(key, entry);
1194 }
1195 }
1196
1197 pub(crate) fn redis_cache_get(&self, ck: &RedisCacheKey) -> Option<CachedRedisValue> {
1198 if !self.redis_global_enabled.load(Ordering::Relaxed) {
1199 return None;
1200 }
1201 let entry = self.fetch_redis_cache(ck)?;
1202 self.redis_cache_hits.fetch_add(1, Ordering::Relaxed);
1203 Some(entry.value)
1204 }
1205
1206 pub(crate) fn redis_cache_put(&self, ck: RedisCacheKey, value: CachedRedisValue) {
1207 self.redis_cache_put_with_ttl(ck, value, 0);
1208 }
1209
1210 pub(crate) fn redis_cache_put_with_ttl(
1211 &self,
1212 ck: RedisCacheKey,
1213 value: CachedRedisValue,
1214 ttl_ms: u64,
1215 ) {
1216 if !self.redis_global_enabled.load(Ordering::Relaxed) {
1217 return;
1218 }
1219 self.redis_cache_misses.fetch_add(1, Ordering::Relaxed);
1220 self.save_redis_cache(
1221 ck,
1222 CachedRedisEntry {
1223 value,
1224 cached_at: Instant::now(),
1225 ttl_ms,
1226 },
1227 );
1228 }
1229
1230 #[allow(dead_code)]
1231 fn clear_redis_cache(&self) {
1232 if let Ok(mut cache) = self.redis_cache.write() {
1233 cache.clear();
1234 }
1235 }
1236
1237 #[inline]
1238 fn result_cache_enabled(&self) -> bool {
1239 self.result_cache_enabled.load(Ordering::Relaxed)
1240 }
1241
1242 #[inline]
1243 fn result_cache_ttl(&self) -> Duration {
1244 Duration::from_millis(self.result_cache_ttl_ms.load(Ordering::Relaxed))
1245 }
1246}
1247
1248pub fn runtime() -> &'static KnowledgeRuntime {
1249 static RUNTIME: OnceLock<KnowledgeRuntime> = OnceLock::new();
1250 RUNTIME.get_or_init(|| KnowledgeRuntime::new(1024))
1251}
1252
1253#[cfg(test)]
1254pub(crate) struct RuntimeTestGuard(tokio::sync::Mutex<()>);
1255
1256#[cfg(test)]
1257impl RuntimeTestGuard {
1258 pub(crate) fn lock(&self) -> Result<tokio::sync::MutexGuard<'_, ()>, std::convert::Infallible> {
1259 Ok(self.0.blocking_lock())
1260 }
1261
1262 pub(crate) async fn lock_async(&self) -> tokio::sync::MutexGuard<'_, ()> {
1263 self.0.lock().await
1264 }
1265}
1266
1267#[cfg(test)]
1268pub(crate) fn runtime_test_guard() -> &'static RuntimeTestGuard {
1269 static GUARD: OnceLock<RuntimeTestGuard> = OnceLock::new();
1270 GUARD.get_or_init(|| RuntimeTestGuard(tokio::sync::Mutex::new(())))
1271}
1272
1273fn result_cache_key(handle: &ProviderHandle, req: &QueryRequest) -> ResultCacheKey {
1274 ResultCacheKey {
1275 datasource_id: handle.datasource_id.clone(),
1276 generation: handle.generation,
1277 query_hash: stable_hash(&req.sql),
1278 params_hash: stable_params_hash(&req.params),
1279 mode: match req.mode {
1280 QueryMode::Many => QueryModeTag::Many,
1281 QueryMode::FirstRow => QueryModeTag::FirstRow,
1282 },
1283 }
1284}
1285
1286fn result_cache_key_fields(
1287 handle: &ProviderHandle,
1288 sql: &str,
1289 params: &[DataField],
1290 mode: QueryModeTag,
1291) -> ResultCacheKey {
1292 ResultCacheKey {
1293 datasource_id: handle.datasource_id.clone(),
1294 generation: handle.generation,
1295 query_hash: stable_hash(sql),
1296 params_hash: stable_field_params_hash(params),
1297 mode,
1298 }
1299}
1300
1301fn query_mode_tag(mode: &QueryMode) -> QueryModeTag {
1302 match mode {
1303 QueryMode::Many => QueryModeTag::Many,
1304 QueryMode::FirstRow => QueryModeTag::FirstRow,
1305 }
1306}
1307
1308fn stable_hash(value: &str) -> u64 {
1309 let mut hasher = DefaultHasher::new();
1310 value.hash(&mut hasher);
1311 hasher.finish()
1312}
1313
1314fn stable_params_hash(params: &[QueryParam]) -> u64 {
1315 let mut hasher = DefaultHasher::new();
1316 for param in params {
1317 param.name.hash(&mut hasher);
1318 match ¶m.value {
1319 QueryValue::Null => 0u8.hash(&mut hasher),
1320 QueryValue::Bool(value) => {
1321 1u8.hash(&mut hasher);
1322 value.hash(&mut hasher);
1323 }
1324 QueryValue::Int(value) => {
1325 2u8.hash(&mut hasher);
1326 value.hash(&mut hasher);
1327 }
1328 QueryValue::Float(value) => {
1329 3u8.hash(&mut hasher);
1330 value.to_bits().hash(&mut hasher);
1331 }
1332 QueryValue::Text(value) => {
1333 4u8.hash(&mut hasher);
1334 value.hash(&mut hasher);
1335 }
1336 }
1337 }
1338 hasher.finish()
1339}
1340
1341fn stable_field_params_hash(params: &[DataField]) -> u64 {
1342 let mut hasher = DefaultHasher::new();
1343 for field in params {
1344 field.get_name().hash(&mut hasher);
1345 match field.get_value() {
1346 Value::Null | Value::Ignore(_) => 0u8.hash(&mut hasher),
1347 Value::Bool(value) => {
1348 1u8.hash(&mut hasher);
1349 value.hash(&mut hasher);
1350 }
1351 Value::Digit(value) => {
1352 2u8.hash(&mut hasher);
1353 value.hash(&mut hasher);
1354 }
1355 Value::Float(value) => {
1356 3u8.hash(&mut hasher);
1357 value.to_bits().hash(&mut hasher);
1358 }
1359 Value::Chars(value) => {
1360 4u8.hash(&mut hasher);
1361 value.hash(&mut hasher);
1362 }
1363 Value::Symbol(value) => {
1364 5u8.hash(&mut hasher);
1365 value.hash(&mut hasher);
1366 }
1367 Value::Time(value) => {
1368 6u8.hash(&mut hasher);
1369 value.hash(&mut hasher);
1370 }
1371 Value::Hex(value) => {
1372 7u8.hash(&mut hasher);
1373 value.to_string().hash(&mut hasher);
1374 }
1375 Value::IpNet(value) => {
1376 8u8.hash(&mut hasher);
1377 value.to_string().hash(&mut hasher);
1378 }
1379 Value::IpAddr(value) => {
1380 9u8.hash(&mut hasher);
1381 value.hash(&mut hasher);
1382 }
1383 Value::Obj(value) => {
1384 10u8.hash(&mut hasher);
1385 format!("{:?}", value).hash(&mut hasher);
1386 }
1387 Value::Array(value) => {
1388 11u8.hash(&mut hasher);
1389 format!("{:?}", value).hash(&mut hasher);
1390 }
1391 Value::Domain(value) => {
1392 12u8.hash(&mut hasher);
1393 value.0.hash(&mut hasher);
1394 }
1395 Value::Url(value) => {
1396 13u8.hash(&mut hasher);
1397 value.0.hash(&mut hasher);
1398 }
1399 Value::Email(value) => {
1400 14u8.hash(&mut hasher);
1401 value.0.hash(&mut hasher);
1402 }
1403 Value::IdCard(value) => {
1404 15u8.hash(&mut hasher);
1405 value.0.hash(&mut hasher);
1406 }
1407 Value::MobilePhone(value) => {
1408 16u8.hash(&mut hasher);
1409 value.0.hash(&mut hasher);
1410 }
1411 Value::BigUint(value) => {
1412 17u8.hash(&mut hasher);
1413 value.to_string().hash(&mut hasher);
1414 }
1415 }
1416 }
1417 hasher.finish()
1418}
1419
1420pub fn fields_to_params(params: &[DataField]) -> Vec<QueryParam> {
1421 params
1422 .iter()
1423 .map(|field| {
1424 let value = match field.get_value() {
1425 Value::Null | Value::Ignore(_) => QueryValue::Null,
1426 Value::Bool(value) => QueryValue::Bool(*value),
1427 Value::Digit(value) => QueryValue::Int(*value),
1428 Value::Float(value) => QueryValue::Float(*value),
1429 Value::Chars(value) => QueryValue::Text(value.to_string()),
1430 Value::Symbol(value) => QueryValue::Text(value.to_string()),
1431 Value::Time(value) => QueryValue::Text(value.to_string()),
1432 Value::Hex(value) => QueryValue::Text(value.to_string()),
1433 Value::IpNet(value) => QueryValue::Text(value.to_string()),
1434 Value::IpAddr(value) => QueryValue::Text(value.to_string()),
1435 Value::Obj(value) => QueryValue::Text(format!("{:?}", value)),
1436 Value::Array(value) => QueryValue::Text(format!("{:?}", value)),
1437 Value::Domain(value) => QueryValue::Text(value.0.to_string()),
1438 Value::Url(value) => QueryValue::Text(value.0.to_string()),
1439 Value::Email(value) => QueryValue::Text(value.0.to_string()),
1440 Value::IdCard(value) => QueryValue::Text(value.0.to_string()),
1441 Value::MobilePhone(value) => QueryValue::Text(value.0.to_string()),
1442 Value::BigUint(value) => QueryValue::Text(value.to_string()),
1444 };
1445 QueryParam {
1446 name: field.get_name().to_string(),
1447 value,
1448 }
1449 })
1450 .collect()
1451}
1452
1453pub fn params_to_fields(params: &[QueryParam]) -> Vec<DataField> {
1454 params
1455 .iter()
1456 .map(|param| match ¶m.value {
1457 QueryValue::Null => {
1458 DataField::new(DataType::default(), param.name.clone(), Value::Null)
1459 }
1460 QueryValue::Bool(value) => {
1461 DataField::new(DataType::default(), param.name.clone(), Value::Bool(*value))
1462 }
1463 QueryValue::Int(value) => DataField::from_digit(param.name.clone(), *value),
1464 QueryValue::Float(value) => DataField::from_float(param.name.clone(), *value),
1465 QueryValue::Text(value) => DataField::from_chars(param.name.clone(), value.clone()),
1466 })
1467 .collect()
1468}
1469
1470#[cfg(test)]
1471mod tests {
1472 use super::*;
1473 use async_trait::async_trait;
1474 use std::sync::Arc;
1475 use wp_model_core::model::Value;
1476
1477 struct TestProvider {
1478 value: &'static str,
1479 }
1480
1481 #[async_trait]
1482 impl ProviderExecutor for TestProvider {
1483 fn query(&self, _sql: &str) -> KnowledgeResult<Vec<RowData>> {
1484 Ok(vec![vec![DataField::from_chars("value", self.value)]])
1485 }
1486
1487 fn query_fields(&self, _sql: &str, _params: &[DataField]) -> KnowledgeResult<Vec<RowData>> {
1488 self.query("")
1489 }
1490
1491 fn query_row(&self, _sql: &str) -> KnowledgeResult<RowData> {
1492 Ok(vec![DataField::from_chars("value", self.value)])
1493 }
1494
1495 fn query_named_fields(
1496 &self,
1497 _sql: &str,
1498 _params: &[DataField],
1499 ) -> KnowledgeResult<RowData> {
1500 self.query_row("")
1501 }
1502 }
1503
1504 #[test]
1505 fn query_param_hash_is_stable() {
1506 let params = vec![
1507 QueryParam {
1508 name: ":id".to_string(),
1509 value: QueryValue::Int(7),
1510 },
1511 QueryParam {
1512 name: ":name".to_string(),
1513 value: QueryValue::Text("abc".to_string()),
1514 },
1515 ];
1516 assert_eq!(stable_params_hash(¶ms), stable_params_hash(¶ms));
1517 }
1518
1519 #[test]
1520 fn fields_to_params_preserves_raw_chars_value() {
1521 let fields = [DataField::from_chars(
1522 ":name".to_string(),
1523 "令狐冲".to_string(),
1524 )];
1525 let params = fields_to_params(&fields);
1526 assert_eq!(params.len(), 1);
1527 match ¶ms[0].value {
1528 QueryValue::Text(value) => assert_eq!(value, "令狐冲"),
1529 other => panic!("unexpected param value: {other:?}"),
1530 }
1531 let roundtrip = params_to_fields(¶ms);
1532 assert!(matches!(roundtrip[0].get_value(), Value::Chars(_)));
1533 }
1534
1535 #[test]
1536 fn fields_to_params_biguint_is_text() {
1537 use num_bigint::BigUint;
1538 use std::str::FromStr;
1539
1540 let fields = [DataField::new(
1541 wp_model_core::model::DataType::BigInt,
1542 ":ip_num",
1543 Value::BigUint(BigUint::from_str("382824323044708348099391746388336347272").unwrap()),
1544 )];
1545 let params = fields_to_params(&fields);
1547 assert_eq!(params.len(), 1);
1548 match ¶ms[0].value {
1549 QueryValue::Text(value) => {
1550 assert_eq!(value, "382824323044708348099391746388336347272")
1551 }
1552 other => panic!("unexpected param value: {other:?}"),
1553 }
1554 assert_eq!(
1556 stable_field_params_hash(&fields),
1557 stable_field_params_hash(&fields)
1558 );
1559 }
1560
1561 #[tokio::test(flavor = "current_thread")]
1562 async fn sqlite_async_bridge_keeps_captured_handle_after_reload() {
1563 let _guard = runtime_test_guard().lock_async().await;
1564 runtime()
1565 .install_provider(
1566 ProviderKind::SqliteAuthority,
1567 DatasourceId("sqlite:old".to_string()),
1568 |_generation| Ok(Arc::new(TestProvider { value: "old" })),
1569 )
1570 .expect("install old provider");
1571 let old_handle = runtime().current_handle().expect("current old handle");
1572
1573 runtime()
1574 .install_provider(
1575 ProviderKind::SqliteAuthority,
1576 DatasourceId("sqlite:new".to_string()),
1577 |_generation| Ok(Arc::new(TestProvider { value: "new" })),
1578 )
1579 .expect("install new provider");
1580
1581 let req = QueryRequest::first_row("SELECT value", Vec::new(), CachePolicy::Bypass);
1582 let row = task::spawn_blocking(move || runtime().execute_with_handle(&old_handle, &req))
1583 .await
1584 .expect("join sqlite bridge")
1585 .expect("execute old handle")
1586 .into_row();
1587 assert_eq!(row[0].to_string(), "chars(old)");
1588 }
1589
1590 #[test]
1591 fn named_providers_are_registered_and_queryable() {
1592 let _guard = runtime_test_guard()
1593 .lock()
1594 .expect("named provider test guard");
1595 runtime()
1596 .install_provider_named(
1597 "geo",
1598 ProviderKind::Postgres,
1599 DatasourceId::from_seed(ProviderKind::Postgres, "geo"),
1600 |_generation| Ok(Arc::new(TestProvider { value: "geo-value" })),
1601 false,
1602 )
1603 .expect("install geo");
1604 runtime()
1605 .install_provider_named(
1606 "asset",
1607 ProviderKind::Postgres,
1608 DatasourceId::from_seed(ProviderKind::Postgres, "asset"),
1609 |_generation| {
1610 Ok(Arc::new(TestProvider {
1611 value: "asset-value",
1612 }))
1613 },
1614 true,
1615 )
1616 .expect("install asset (default)");
1617
1618 assert!(runtime().provider_exists("geo"));
1619 assert!(runtime().provider_exists("asset"));
1620 assert!(!runtime().provider_exists("nope"));
1621 let names = runtime().provider_names();
1622 assert!(names.contains(&"geo".to_string()));
1623 assert!(names.contains(&"asset".to_string()));
1624
1625 let req = QueryRequest::first_row("SELECT value", Vec::new(), CachePolicy::Bypass);
1626 let geo_row = runtime()
1627 .execute_for("geo", &req)
1628 .expect("execute geo")
1629 .into_row();
1630 assert_eq!(geo_row[0].to_string(), "chars(geo-value)");
1631
1632 let row = runtime()
1633 .execute_first_row_fields_for("asset", "SELECT value", &[], CachePolicy::Bypass)
1634 .expect("execute asset");
1635 assert_eq!(row[0].to_string(), "chars(asset-value)");
1636
1637 let default_row = runtime()
1638 .execute(&req)
1639 .expect("execute default (asset)")
1640 .into_row();
1641 assert_eq!(default_row[0].to_string(), "chars(asset-value)");
1642 }
1643
1644 #[test]
1645 fn prune_named_providers_except_keeps_only_selected() {
1646 let _guard = runtime_test_guard().lock().expect("prune test guard");
1647 runtime()
1648 .install_provider_named(
1649 "a",
1650 ProviderKind::Postgres,
1651 DatasourceId::from_seed(ProviderKind::Postgres, "a"),
1652 |_generation| Ok(Arc::new(TestProvider { value: "a" })),
1653 false,
1654 )
1655 .expect("install a");
1656 runtime()
1657 .install_provider_named(
1658 "b",
1659 ProviderKind::Postgres,
1660 DatasourceId::from_seed(ProviderKind::Postgres, "b"),
1661 |_generation| Ok(Arc::new(TestProvider { value: "b" })),
1662 false,
1663 )
1664 .expect("install b");
1665
1666 let keep: HashSet<String> = ["a".to_string()].into_iter().collect();
1667 runtime().prune_named_providers_except(&keep);
1668 assert!(runtime().provider_exists("a"));
1669 assert!(!runtime().provider_exists("b"));
1670 }
1671
1672 #[test]
1673 fn provider_by_name_unknown_returns_error() {
1674 let err = runtime()
1675 .provider_by_name("nope")
1676 .err()
1677 .expect("unknown provider name should error");
1678 assert!(err.to_string().contains("not initialized"));
1679 }
1680
1681 #[test]
1682 fn execute_first_row_fields_for_unknown_returns_error() {
1683 let err = runtime()
1684 .execute_first_row_fields_for("nope", "SELECT value", &[], CachePolicy::Bypass)
1685 .expect_err("unknown provider name");
1686 assert!(err.to_string().contains("not initialized"));
1687 }
1688
1689 #[tokio::test(flavor = "current_thread")]
1690 async fn named_provider_async_dispatch_works() {
1691 let _guard = runtime_test_guard().lock_async().await;
1692 runtime()
1693 .install_provider_named(
1694 "geo",
1695 ProviderKind::Postgres,
1696 DatasourceId::from_seed(ProviderKind::Postgres, "geo"),
1697 |_generation| Ok(Arc::new(TestProvider { value: "geo-async" })),
1698 false,
1699 )
1700 .expect("install geo");
1701
1702 let req = QueryRequest::first_row("SELECT value", Vec::new(), CachePolicy::Bypass);
1703 let row = runtime()
1704 .execute_for_async("geo", &req)
1705 .await
1706 .expect("async geo")
1707 .into_row();
1708 assert_eq!(row[0].to_string(), "chars(geo-async)");
1709
1710 let row = runtime()
1711 .execute_first_row_fields_for_async("geo", "SELECT value", &[], CachePolicy::Bypass)
1712 .await
1713 .expect("async geo first row");
1714 assert_eq!(row[0].to_string(), "chars(geo-async)");
1715 }
1716
1717 fn redis_ck(cmd: RedisCmdTag, generation: u64, key: &str, args: &[&str]) -> RedisCacheKey {
1722 let mut hasher = DefaultHasher::new();
1723 key.hash(&mut hasher);
1724 let key_hash = hasher.finish();
1725 let mut hasher = DefaultHasher::new();
1726 for arg in args {
1727 arg.hash(&mut hasher);
1728 }
1729 let args_hash = hasher.finish();
1730 RedisCacheKey {
1731 generation,
1732 cmd_tag: cmd,
1733 key_hash,
1734 args_hash,
1735 }
1736 }
1737
1738 #[test]
1739 fn redis_cache_hit_and_miss() {
1740 let rt = KnowledgeRuntime::new(64);
1741 rt.configure_redis_cache(true, 64);
1742
1743 let ck = redis_ck(RedisCmdTag::Get, 1, "user:1", &[]);
1744 assert!(rt.redis_cache_get(&ck).is_none());
1746 rt.redis_cache_put(ck.clone(), CachedRedisValue::Bool(true));
1748 let val = rt.redis_cache_get(&ck).expect("should hit cache");
1750 assert!(matches!(val, CachedRedisValue::Bool(true)));
1751 }
1752
1753 #[test]
1754 fn redis_cache_global_enabled_access() {
1755 let rt = KnowledgeRuntime::new(64);
1756 rt.configure_redis_cache(true, 64);
1757
1758 let ck = redis_ck(RedisCmdTag::Get, 1, "k", &[]);
1759 rt.redis_cache_put(ck.clone(), CachedRedisValue::Bool(true));
1760 assert!(rt.redis_cache_get(&ck).is_some());
1761 }
1762
1763 #[test]
1764 fn redis_cache_global_disabled_blocks_all() {
1765 let rt = KnowledgeRuntime::new(64);
1766 rt.configure_redis_cache(false, 64);
1767
1768 let ck = redis_ck(RedisCmdTag::BfExists, 1, "any_key", &["item"]);
1769 rt.redis_cache_put(ck.clone(), CachedRedisValue::Bool(true));
1770 assert!(rt.redis_cache_get(&ck).is_none());
1772 }
1773
1774 #[test]
1775 fn redis_cache_ttl_expiry() {
1776 let rt = KnowledgeRuntime::new(64);
1777 rt.configure_redis_cache(true, 64);
1778
1779 let ck = redis_ck(RedisCmdTag::BfExists, 1, "k", &["item"]);
1780 rt.redis_cache_put_with_ttl(
1782 ck.clone(),
1783 CachedRedisValue::Bool(true),
1784 1, );
1786 assert!(rt.redis_cache_get(&ck).is_some());
1788 std::thread::sleep(std::time::Duration::from_millis(5));
1790 assert!(rt.redis_cache_get(&ck).is_none());
1792 }
1793
1794 #[test]
1795 fn redis_cache_no_ttl_never_expires() {
1796 let rt = KnowledgeRuntime::new(64);
1797 rt.configure_redis_cache(true, 64);
1798
1799 let ck = redis_ck(RedisCmdTag::Get, 1, "k", &[]);
1800 rt.redis_cache_put(ck.clone(), CachedRedisValue::Bool(true));
1802 assert!(rt.redis_cache_get(&ck).is_some());
1804 std::thread::sleep(std::time::Duration::from_millis(5));
1805 assert!(rt.redis_cache_get(&ck).is_some());
1807 }
1808
1809 #[test]
1810 fn redis_cache_generation_isolation() {
1811 let rt = KnowledgeRuntime::new(64);
1812 rt.configure_redis_cache(true, 64);
1813
1814 let ck_gen1 = redis_ck(RedisCmdTag::BfExists, 1, "key", &["item"]);
1815 let ck_gen2 = redis_ck(RedisCmdTag::BfExists, 2, "key", &["item"]);
1816
1817 rt.redis_cache_put(ck_gen1.clone(), CachedRedisValue::Bool(false));
1819
1820 assert!(rt.redis_cache_get(&ck_gen2).is_none());
1822 assert!(rt.redis_cache_get(&ck_gen1).is_some());
1824 }
1825}