use crate::schemasync::PreservationMode;
use bon::Builder;
use serde::{Deserialize, Serialize};
use std::collections::BTreeMap;
use tracing::{debug, trace};
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct PluginConfig {
pub path: String,
#[serde(default)]
pub params: BTreeMap<String, String>,
}
#[derive(Debug, Clone, Deserialize, Serialize, Builder)]
#[serde(deny_unknown_fields)]
pub struct SchemasyncConfig {
pub database: DatabaseConfig,
pub should_generate_mocks: bool,
#[serde(default)]
pub mock_gen_config: SchemasyncMockGenConfig,
#[serde(default)]
#[builder(default)]
pub plugins: BTreeMap<String, PluginConfig>,
#[serde(default)]
#[builder(default)]
pub lint: LintConfig,
}
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct LintConfig {
#[serde(default)]
pub silence_unverifiable_annotations: bool,
}
#[derive(Debug, Clone, Deserialize, Serialize, Default, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum DatabaseProvider {
#[default]
Surrealdb,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct DatabaseConfig {
#[serde(default)]
pub provider: DatabaseProvider,
pub url: String,
#[serde(default)]
pub namespace: String,
#[serde(default)]
pub database: String,
#[serde(default = "default_timeout")]
pub timeout: u64,
#[serde(default)]
pub accesses: AccessesSource,
#[serde(default)]
pub functions: Option<FunctionsSource>,
#[serde(default)]
pub analyzers: Option<AnalyzersSource>,
#[serde(skip)]
pub resolved: ResolvedDatabaseItems,
}
fn default_timeout() -> u64 {
60
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct AccessConfig {
pub name: String,
pub access_type: AccessType,
pub table_name: String,
}
#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
pub enum AccessType {
System,
Record,
Bearer,
Jwt,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
#[serde(untagged)]
pub enum AccessesSource {
Inline(Vec<AccessConfig>),
Path { path: String },
}
impl Default for AccessesSource {
fn default() -> Self {
AccessesSource::Inline(vec![])
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct FunctionsSource {
pub path: String,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct AnalyzersSource {
pub path: String,
}
#[derive(Debug, Clone, Default)]
pub struct ResolvedDatabaseItems {
pub access_surql: Option<String>,
pub functions_surql: Option<String>,
pub analyzers_surql: Option<String>,
}
#[derive(Debug, Clone, Deserialize, Serialize, Builder)]
#[serde(deny_unknown_fields)]
#[serde(default)]
pub struct SchemasyncMockGenConfig {
pub default_record_count: usize,
pub default_preservation_mode: PreservationMode,
pub full_refresh_mode: bool,
#[builder(default = true)]
pub scripting_asserts: bool,
}
impl Default for SchemasyncMockGenConfig {
fn default() -> Self {
Self {
default_record_count: 10,
default_preservation_mode: PreservationMode::default(),
full_refresh_mode: false,
scripting_asserts: true,
}
}
}
impl Default for DatabaseConfig {
fn default() -> Self {
Self {
provider: DatabaseProvider::default(),
url: String::new(),
namespace: String::new(),
database: String::new(),
timeout: default_timeout(),
accesses: AccessesSource::default(),
functions: None,
analyzers: None,
resolved: ResolvedDatabaseItems::default(),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct MockOverrides {
pub skip_mocks: bool,
pub full_refresh: bool,
}
impl SchemasyncConfig {
pub fn apply_mock_overrides(&mut self, overrides: &MockOverrides) {
if overrides.skip_mocks {
self.should_generate_mocks = false;
}
if overrides.full_refresh {
self.mock_gen_config.full_refresh_mode = true;
}
}
}
#[derive(Debug, Clone, Default)]
pub struct ConnectionOverrides {
pub url: Option<String>,
pub namespace: Option<String>,
pub database: Option<String>,
}
impl DatabaseConfig {
pub fn apply_connection_overrides(&mut self, overrides: &ConnectionOverrides) {
if let Some(url) = &overrides.url {
self.url = url.clone();
}
if let Some(namespace) = &overrides.namespace {
self.namespace = namespace.clone();
}
if let Some(database) = &overrides.database {
self.database = database.clone();
}
}
pub fn unresolved_connection_var(&self) -> Option<String> {
[&self.url, &self.namespace, &self.database]
.into_iter()
.find_map(|value| {
let start = value.find("${")? + 2;
let name: String = value[start..]
.chars()
.take_while(|c| c.is_ascii_uppercase() || c.is_ascii_digit() || *c == '_')
.collect();
(!name.is_empty()).then_some(name)
})
}
pub fn for_testing() -> Self {
debug!("Creating database configuration for testing environment");
let config = Self {
provider: DatabaseProvider::Surrealdb,
url: "http://localhost:8000".to_string(),
namespace: "test".to_string(),
database: "test".to_string(),
accesses: AccessesSource::Inline(vec![AccessConfig {
name: "user".to_owned(),
access_type: AccessType::Record,
table_name: "user".to_owned(),
}]),
functions: None,
analyzers: None,
resolved: ResolvedDatabaseItems::default(),
timeout: 60,
};
trace!(
"Test database config - URL: {}, namespace: {}, database: {}, timeout: {}s",
config.url, config.namespace, config.database, config.timeout
);
if let AccessesSource::Inline(ref accesses) = config.accesses {
trace!("Test access configs: {} entries", accesses.len());
}
config
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mock_overrides_only_switch_settings_on() {
let mut config: SchemasyncConfig =
toml::from_str("should_generate_mocks = true\n[database]\nurl = \"x\"\n").unwrap();
config.apply_mock_overrides(&MockOverrides::default());
assert!(config.should_generate_mocks);
assert!(!config.mock_gen_config.full_refresh_mode);
config.apply_mock_overrides(&MockOverrides {
skip_mocks: true,
full_refresh: true,
});
assert!(!config.should_generate_mocks);
assert!(config.mock_gen_config.full_refresh_mode);
}
#[test]
fn connection_overrides_replace_unresolved_settings() {
let mut database = DatabaseConfig::for_testing();
database.url = "${SURREALDB_URL}".to_string();
database.namespace = "prefix_${SURREALDB_NS}".to_string();
assert_eq!(
database.unresolved_connection_var().as_deref(),
Some("SURREALDB_URL")
);
database.apply_connection_overrides(&ConnectionOverrides {
url: Some("http://localhost:8000".to_string()),
namespace: None,
database: None,
});
assert_eq!(database.url, "http://localhost:8000");
assert_eq!(
database.unresolved_connection_var().as_deref(),
Some("SURREALDB_NS")
);
database.apply_connection_overrides(&ConnectionOverrides {
url: None,
namespace: Some("app".to_string()),
database: Some("main".to_string()),
});
assert_eq!(database.unresolved_connection_var(), None);
assert_eq!(database.database, "main");
}
}