1use crate::handle::DbPool;
2use crate::migrate_cli::run_migrate;
3use crate::tx::inject_conn;
4use sova_core::extend::StateMap;
5use sova_core::{App, Error, Plugin};
6use sea_orm::{ConnectOptions, Database, DatabaseConnection};
7use sea_orm_migration::MigratorTrait;
8use std::future::Future;
9use std::pin::Pin;
10use std::sync::Arc;
11
12type MigrateFn = Arc<
13 dyn Fn(
14 DatabaseConnection,
15 Vec<String>,
16 ) -> Pin<Box<dyn Future<Output = Result<(), Error>> + Send>>
17 + Send
18 + Sync,
19>;
20
21type SeedFn = Arc<
22 dyn Fn(Arc<StateMap>) -> Pin<Box<dyn Future<Output = Result<(), Error>> + Send>>
23 + Send
24 + Sync,
25>;
26
27pub struct Db {
29 url: String,
30 url_pinned: bool,
32 sqlx_logging: bool,
34 migrate: Option<MigrateFn>,
35 seed: Option<SeedFn>,
36}
37
38impl Db {
39 pub fn from_env() -> Self {
40 let url = std::env::var("DATABASE_URL").unwrap_or_default();
41 Self {
42 url,
43 url_pinned: false,
44 sqlx_logging: false,
45 migrate: None,
46 seed: None,
47 }
48 }
49
50 pub fn url(mut self, url: impl Into<String>) -> Self {
52 self.url = url.into();
53 self.url_pinned = true;
54 self
55 }
56
57 pub fn sqlx_logging(mut self, on: bool) -> Self {
59 self.sqlx_logging = on;
60 self
61 }
62
63 pub fn migrations<M: MigratorTrait + 'static>(mut self) -> Self {
65 self.migrate = Some(Arc::new(move |conn, args| {
66 Box::pin(async move { run_migrate::<M>(conn, &args).await })
67 }));
68 self
69 }
70
71 pub fn seed<F, Fut, E>(mut self, f: F) -> Self
75 where
76 F: Fn(Arc<StateMap>) -> Fut + Send + Sync + 'static,
77 Fut: Future<Output = Result<(), E>> + Send + 'static,
78 E: Into<Error> + Send + 'static,
79 {
80 self.seed = Some(Arc::new(move |state| {
81 let fut = f(state);
82 Box::pin(async move { fut.await.map_err(Into::into) })
83 }));
84 self
85 }
86}
87
88impl Plugin for Db {
89 fn id(&self) -> &'static str {
90 "db"
91 }
92
93 fn meta(&self) -> sova_core::PluginMeta {
94 sova_core::PluginMeta::new("Database")
95 .description("SeaORM pool, migrate CLI, optional seed CLI")
96 .version(env!("CARGO_PKG_VERSION"))
97 }
98
99 fn install(mut self, app: &mut App) {
100 if !self.url_pinned {
102 if let Ok(u) = std::env::var("DATABASE_URL") {
103 if !u.is_empty() {
104 self.url = u;
105 }
106 }
107 if self.url.is_empty() {
108 if let Some(u) = app
109 .config_doc()
110 .and_then(|d| d.section("db"))
111 .and_then(|s| s.get("url").and_then(|v| v.as_str()).map(str::to_string))
112 {
113 self.url = u;
114 }
115 }
116 }
117
118 if self.url.is_empty() {
119 app.on_startup(|_state| async {
120 Err(Error::Internal(
121 "database url is empty; set DATABASE_URL or [db] url in sova.toml".into(),
122 ))
123 });
124 return;
125 }
126
127 let pool = DbPool::new();
128 app.state(pool.clone());
129
130 let url = self.url.clone();
131 let pool_start = pool.clone();
132 let sqlx_logging = self.sqlx_logging;
133 app.on_startup(move |_state| {
134 let url = url.clone();
135 let pool = pool_start.clone();
136 async move {
137 let mut opt = ConnectOptions::new(url);
138 opt.sqlx_logging(sqlx_logging);
139 let conn = Database::connect(opt)
140 .await
141 .map_err(|e| Error::Internal(format!("db connect: {e}")))?;
142 conn.ping()
143 .await
144 .map_err(|e| Error::Internal(format!("db ping: {e}")))?;
145 pool.set(conn);
146 Ok(())
147 }
148 });
149
150 let pool_stop = pool.clone();
151 app.on_shutdown(move || {
152 let pool = pool_stop.clone();
153 async move {
154 pool.clear();
155 }
156 });
157
158 app.use_middleware(inject_conn(pool.clone()));
159
160 let pool_check = pool.clone();
161 app.register_check("db", move |_state| {
162 let pool = pool_check.clone();
163 async move {
164 let conn = pool.get().map_err(Error::from)?;
165 conn.ping()
166 .await
167 .map_err(|e| Error::Internal(format!("db ping: {e}")))?;
168 Ok(())
169 }
170 });
171
172 if let Some(migrate) = self.migrate {
173 let pool_cli = pool.clone();
174 app.register_cli("migrate", move |_state, args| {
175 let pool = pool_cli.clone();
176 let migrate = Arc::clone(&migrate);
177 async move {
178 let conn = pool.get().map_err(Error::from)?;
179 migrate(conn, args).await
180 }
181 });
182 }
183
184 if let Some(seed) = self.seed {
185 app.register_cli("seed", move |state, _args| {
186 let seed = Arc::clone(&seed);
187 async move { seed(state).await }
188 });
189 }
190 }
191}