use std::sync::Arc;
use std::time::Duration;
use thiserror::Error;
use crate::cache::Cache;
const DEFAULT_KEY_PREFIX: &str = "schema_cache";
#[derive(Debug, Error)]
pub enum SchemaCacheError {
#[error("缓存操作失败: {0}")]
CacheError(#[from] sz_orm_core::CacheError),
#[error("Schema 加载失败: {0}")]
LoaderError(String),
#[error("Schema 序列化失败: {0}")]
Serialize(String),
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct ColumnDefinition {
pub name: String,
pub data_type: String,
pub nullable: bool,
pub primary_key: bool,
pub auto_increment: bool,
pub unsigned: bool,
pub default: Option<String>,
pub comment: Option<String>,
}
impl ColumnDefinition {
pub fn new(name: impl Into<String>, data_type: impl Into<String>) -> Self {
Self {
name: name.into(),
data_type: data_type.into(),
nullable: true,
primary_key: false,
auto_increment: false,
unsigned: false,
default: None,
comment: None,
}
}
pub fn nullable(mut self, nullable: bool) -> Self {
self.nullable = nullable;
self
}
pub fn primary_key(mut self, primary_key: bool) -> Self {
self.primary_key = primary_key;
self
}
pub fn auto_increment(mut self, auto_increment: bool) -> Self {
self.auto_increment = auto_increment;
self
}
pub fn unsigned(mut self, unsigned: bool) -> Self {
self.unsigned = unsigned;
self
}
pub fn default_value(mut self, default: impl Into<String>) -> Self {
self.default = Some(default.into());
self
}
pub fn comment(mut self, comment: impl Into<String>) -> Self {
self.comment = Some(comment.into());
self
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct TableSchema {
pub table_name: String,
pub columns: Vec<ColumnDefinition>,
pub primary_keys: Vec<String>,
pub cached_at: i64,
}
impl TableSchema {
pub fn new(table_name: impl Into<String>, columns: Vec<ColumnDefinition>) -> Self {
let table_name = table_name.into();
let primary_keys: Vec<String> = columns
.iter()
.filter(|col| col.primary_key)
.map(|col| col.name.clone())
.collect();
Self {
table_name,
columns,
primary_keys,
cached_at: chrono::Utc::now().timestamp(),
}
}
pub fn column(&self, name: &str) -> Option<&ColumnDefinition> {
self.columns.iter().find(|col| col.name == name)
}
pub fn has_column(&self, name: &str) -> bool {
self.column(name).is_some()
}
pub fn column_names(&self) -> Vec<&str> {
self.columns.iter().map(|col| col.name.as_str()).collect()
}
}
pub struct SchemaCache {
cache: Arc<Cache>,
key_prefix: String,
tag_name: String,
}
impl SchemaCache {
pub fn new(cache: Arc<Cache>) -> Self {
Self {
cache,
key_prefix: DEFAULT_KEY_PREFIX.to_string(),
tag_name: DEFAULT_KEY_PREFIX.to_string(),
}
}
pub fn with_prefix(cache: Arc<Cache>, prefix: impl Into<String>) -> Self {
let prefix = prefix.into();
Self {
cache,
tag_name: prefix.clone(),
key_prefix: prefix,
}
}
pub fn cache_key(&self, table_name: &str) -> String {
format!("{}:{}", self.key_prefix, table_name)
}
pub fn get_schema(&self, table_name: &str) -> Result<Option<TableSchema>, SchemaCacheError> {
let key = self.cache_key(table_name);
let result = self.cache.get::<TableSchema>(&key)?;
Ok(result)
}
pub fn remember_schema<F>(
&self,
table_name: &str,
loader: F,
) -> Result<TableSchema, SchemaCacheError>
where
F: FnOnce(&str) -> Result<TableSchema, SchemaCacheError>,
{
if let Some(schema) = self.get_schema(table_name)? {
return Ok(schema);
}
let schema = loader(table_name)?;
self.set_schema(table_name, &schema, None)?;
Ok(schema)
}
pub fn set_schema(
&self,
table_name: &str,
schema: &TableSchema,
ttl: Option<Duration>,
) -> Result<(), SchemaCacheError> {
let key = self.cache_key(table_name);
self.cache.tag(&self.tag_name).set(&key, schema, ttl)?;
Ok(())
}
pub fn forget_schema(&self, table_name: &str) -> Result<(), SchemaCacheError> {
let key = self.cache_key(table_name);
self.cache.delete(&key)?;
Ok(())
}
pub fn clear(&self) -> Result<(), SchemaCacheError> {
self.cache.tag(&self.tag_name).clear()?;
Ok(())
}
pub fn has_schema(&self, table_name: &str) -> Result<bool, SchemaCacheError> {
let key = self.cache_key(table_name);
Ok(self.cache.has(&key)?)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::cache::{Cache, MemoryCacheDriver};
use std::sync::atomic::{AtomicBool, Ordering};
fn make_cache() -> Arc<Cache> {
let cache = Arc::new(Cache::new());
cache.register_default(MemoryCacheDriver::new());
cache
}
fn make_columns() -> Vec<ColumnDefinition> {
vec![
ColumnDefinition::new("id", "int(11) unsigned")
.nullable(false)
.primary_key(true)
.auto_increment(true)
.unsigned(true),
ColumnDefinition::new("name", "varchar(255)")
.nullable(false)
.default_value(""),
ColumnDefinition::new("email", "varchar(255)")
.nullable(true)
.comment("用户邮箱"),
]
}
#[test]
fn test_column_definition_builder() {
let col = ColumnDefinition::new("id", "int(11) unsigned")
.nullable(false)
.primary_key(true)
.auto_increment(true)
.unsigned(true)
.default_value("0")
.comment("主键");
assert_eq!(col.name, "id");
assert_eq!(col.data_type, "int(11) unsigned");
assert!(!col.nullable);
assert!(col.primary_key);
assert!(col.auto_increment);
assert!(col.unsigned);
assert_eq!(col.default, Some("0".to_string()));
assert_eq!(col.comment, Some("主键".to_string()));
}
#[test]
fn test_column_definition_defaults() {
let col = ColumnDefinition::new("name", "varchar(255)");
assert_eq!(col.name, "name");
assert!(col.nullable);
assert!(!col.primary_key);
assert!(!col.auto_increment);
assert!(!col.unsigned);
assert_eq!(col.default, None);
assert_eq!(col.comment, None);
}
#[test]
fn test_table_schema_primary_keys() {
let schema = TableSchema::new("users", make_columns());
assert_eq!(schema.table_name, "users");
assert_eq!(schema.columns.len(), 3);
assert_eq!(schema.primary_keys, vec!["id"]);
}
#[test]
fn test_table_schema_column_lookup() {
let schema = TableSchema::new("users", make_columns());
assert!(schema.has_column("id"));
assert!(schema.has_column("name"));
assert!(schema.has_column("email"));
assert!(!schema.has_column("nonexistent"));
let col = schema.column("email").unwrap();
assert_eq!(col.data_type, "varchar(255)");
assert_eq!(col.comment, Some("用户邮箱".to_string()));
}
#[test]
fn test_table_schema_column_names() {
let schema = TableSchema::new("users", make_columns());
let names = schema.column_names();
assert_eq!(names, vec!["id", "name", "email"]);
}
#[test]
fn test_table_schema_serde() {
let schema = TableSchema::new("users", make_columns());
let json = serde_json::to_string(&schema).unwrap();
let deserialized: TableSchema = serde_json::from_str(&json).unwrap();
assert_eq!(schema, deserialized);
}
#[test]
fn test_set_get_schema() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache);
let schema = TableSchema::new("users", make_columns());
schema_cache.set_schema("users", &schema, None).unwrap();
let result = schema_cache.get_schema("users").unwrap();
assert!(result.is_some());
let cached = result.unwrap();
assert_eq!(cached.table_name, "users");
assert_eq!(cached.columns.len(), 3);
assert_eq!(cached.primary_keys, vec!["id"]);
assert_eq!(cached.columns[0].name, "id");
assert!(cached.columns[0].primary_key);
assert!(cached.columns[0].auto_increment);
assert!(!cached.columns[0].nullable);
}
#[test]
fn test_get_schema_miss() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache);
let result = schema_cache.get_schema("nonexistent").unwrap();
assert!(result.is_none());
}
#[test]
fn test_remember_schema_cache_hit() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache);
let schema = TableSchema::new("users", make_columns());
schema_cache.set_schema("users", &schema, None).unwrap();
let loader_called = Arc::new(AtomicBool::new(false));
let loader_called_clone = loader_called.clone();
let result = schema_cache
.remember_schema("users", |_| {
loader_called_clone.store(true, Ordering::SeqCst);
Ok(TableSchema::new("users", vec![]))
})
.unwrap();
assert!(
!loader_called.load(Ordering::SeqCst),
"loader should not be called on cache hit"
);
assert_eq!(result.columns.len(), 3);
}
#[test]
fn test_remember_schema_cache_miss() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache.clone());
let result = schema_cache
.remember_schema("orders", |table| {
assert_eq!(table, "orders");
Ok(TableSchema::new(
"orders",
vec![ColumnDefinition::new("id", "bigint(20)")
.primary_key(true)
.auto_increment(true)],
))
})
.unwrap();
assert_eq!(result.table_name, "orders");
assert_eq!(result.columns.len(), 1);
let cached = schema_cache.get_schema("orders").unwrap().unwrap();
assert_eq!(cached.columns.len(), 1);
assert_eq!(cached.columns[0].name, "id");
}
#[test]
fn test_remember_schema_loader_error() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache);
let result = schema_cache.remember_schema("broken", |_| {
Err(SchemaCacheError::LoaderError("数据库连接失败".to_string()))
});
assert!(result.is_err());
match result.unwrap_err() {
SchemaCacheError::LoaderError(msg) => assert_eq!(msg, "数据库连接失败"),
other => panic!("期望 LoaderError,实际: {:?}", other),
}
assert!(schema_cache.get_schema("broken").unwrap().is_none());
}
#[test]
fn test_forget_schema() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache);
let schema = TableSchema::new("users", make_columns());
schema_cache.set_schema("users", &schema, None).unwrap();
assert!(schema_cache.has_schema("users").unwrap());
schema_cache.forget_schema("users").unwrap();
assert!(!schema_cache.has_schema("users").unwrap());
assert!(schema_cache.get_schema("users").unwrap().is_none());
}
#[test]
fn test_clear() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache);
schema_cache
.set_schema("users", &TableSchema::new("users", make_columns()), None)
.unwrap();
schema_cache
.set_schema("orders", &TableSchema::new("orders", make_columns()), None)
.unwrap();
schema_cache
.set_schema(
"products",
&TableSchema::new("products", make_columns()),
None,
)
.unwrap();
assert!(schema_cache.has_schema("users").unwrap());
assert!(schema_cache.has_schema("orders").unwrap());
assert!(schema_cache.has_schema("products").unwrap());
schema_cache.clear().unwrap();
assert!(!schema_cache.has_schema("users").unwrap());
assert!(!schema_cache.has_schema("orders").unwrap());
assert!(!schema_cache.has_schema("products").unwrap());
}
#[test]
fn test_clear_preserves_other_caches() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache.clone());
schema_cache
.set_schema("users", &TableSchema::new("users", make_columns()), None)
.unwrap();
cache.set("business:config", "value", None).unwrap();
schema_cache.clear().unwrap();
assert!(!schema_cache.has_schema("users").unwrap());
assert!(cache.has("business:config").unwrap());
}
#[test]
fn test_has_schema() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache);
assert!(!schema_cache.has_schema("users").unwrap());
let schema = TableSchema::new("users", make_columns());
schema_cache.set_schema("users", &schema, None).unwrap();
assert!(schema_cache.has_schema("users").unwrap());
assert!(!schema_cache.has_schema("orders").unwrap());
}
#[test]
fn test_different_tables_no_conflict() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache);
let users_schema = TableSchema::new(
"users",
vec![
ColumnDefinition::new("id", "int(11)").primary_key(true),
ColumnDefinition::new("name", "varchar(255)"),
],
);
let orders_schema = TableSchema::new(
"orders",
vec![
ColumnDefinition::new("order_id", "bigint(20)").primary_key(true),
ColumnDefinition::new("amount", "decimal(10,2)"),
],
);
schema_cache
.set_schema("users", &users_schema, None)
.unwrap();
schema_cache
.set_schema("orders", &orders_schema, None)
.unwrap();
let users = schema_cache.get_schema("users").unwrap().unwrap();
let orders = schema_cache.get_schema("orders").unwrap().unwrap();
assert_eq!(users.table_name, "users");
assert_eq!(users.columns[0].name, "id");
assert_eq!(orders.table_name, "orders");
assert_eq!(orders.columns[0].name, "order_id");
schema_cache.forget_schema("users").unwrap();
assert!(schema_cache.get_schema("users").unwrap().is_none());
assert!(schema_cache.get_schema("orders").unwrap().is_some());
}
#[test]
fn test_ttl_expiry() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache);
let schema = TableSchema::new("temp_table", make_columns());
schema_cache
.set_schema("temp_table", &schema, Some(Duration::from_millis(50)))
.unwrap();
assert!(schema_cache.has_schema("temp_table").unwrap());
assert!(schema_cache.get_schema("temp_table").unwrap().is_some());
std::thread::sleep(Duration::from_millis(100));
assert!(!schema_cache.has_schema("temp_table").unwrap());
assert!(schema_cache.get_schema("temp_table").unwrap().is_none());
}
#[test]
fn test_default_ttl_no_expiry() {
let cache = make_cache();
let schema_cache = SchemaCache::new(cache);
let schema = TableSchema::new("permanent", make_columns());
schema_cache.set_schema("permanent", &schema, None).unwrap();
std::thread::sleep(Duration::from_millis(50));
assert!(schema_cache.has_schema("permanent").unwrap());
assert!(schema_cache.get_schema("permanent").unwrap().is_some());
}
#[test]
fn test_custom_prefix() {
let cache = make_cache();
let schema_cache = SchemaCache::with_prefix(cache.clone(), "my_schema");
let schema = TableSchema::new("users", make_columns());
schema_cache.set_schema("users", &schema, None).unwrap();
let key = schema_cache.cache_key("users");
assert_eq!(key, "my_schema:users");
assert!(cache.has("my_schema:users").unwrap());
assert!(!cache.has("schema_cache:users").unwrap());
schema_cache.clear().unwrap();
assert!(!cache.has("my_schema:users").unwrap());
}
#[test]
fn test_multiple_prefixes_no_conflict() {
let cache = make_cache();
let schema_cache_1 = SchemaCache::new(cache.clone());
let schema_cache_2 = SchemaCache::with_prefix(cache.clone(), "db2_schema");
let schema = TableSchema::new("users", make_columns());
schema_cache_1.set_schema("users", &schema, None).unwrap();
schema_cache_2.set_schema("users", &schema, None).unwrap();
assert!(schema_cache_1.has_schema("users").unwrap());
assert!(schema_cache_2.has_schema("users").unwrap());
assert!(cache.has("schema_cache:users").unwrap());
assert!(cache.has("db2_schema:users").unwrap());
schema_cache_1.clear().unwrap();
assert!(!schema_cache_1.has_schema("users").unwrap());
assert!(schema_cache_2.has_schema("users").unwrap());
}
#[test]
fn test_cache_key_format() {
let cache = make_cache();
let sc1 = SchemaCache::new(cache.clone());
assert_eq!(sc1.cache_key("users"), "schema_cache:users");
assert_eq!(sc1.cache_key("orders"), "schema_cache:orders");
let sc2 = SchemaCache::with_prefix(cache, "custom_prefix");
assert_eq!(sc2.cache_key("users"), "custom_prefix:users");
}
}