Skip to main content

rp_supabase_codegen/
lib.rs

1#![cfg_attr(all(doc, not(doctest)), doc = include_str!("../README.md"))]
2
3extern crate alloc;
4
5#[cfg(feature = "database")]
6mod database;
7mod emitter;
8pub mod model;
9
10use alloc::collections::{BTreeMap, BTreeSet};
11use std::path::{Path, PathBuf};
12
13use model::{SNAPSHOT_VERSION, Snapshot};
14
15/// Generation or schema acquisition failed.
16#[expect(
17    clippy::error_impl_error,
18    reason = "The module-qualified codegen::Error is unambiguous for callers."
19)]
20#[derive(Debug, thiserror::Error)]
21pub enum Error {
22    #[error("invalid schema generation input: {0}")]
23    Invalid(String),
24    #[error("schema generation I/O failed: {0}")]
25    Io(#[from] std::io::Error),
26    #[error("invalid schema snapshot: {0}")]
27    Json(#[from] serde_json::Error),
28    #[error("PostgreSQL schema introspection failed: {0}")]
29    Database(String),
30}
31
32#[derive(Debug, Clone)]
33struct Config {
34    pub schemas: Vec<String>,
35    pub prelude: String,
36    pub derives: Vec<String>,
37    pub attributes: Vec<String>,
38    pub type_overrides: BTreeMap<String, String>,
39    pub type_attributes: BTreeMap<String, Vec<String>>,
40    pub runtime_path: String,
41    pub not_null: BTreeMap<String, Vec<String>>,
42    pub column_types: BTreeMap<String, String>,
43    pub json_types: BTreeMap<String, String>,
44    pub reexport_macros: bool,
45    pub relationship_aliases: BTreeMap<String, String>,
46    pub strict_args: bool,
47    pub strict_functions: BTreeSet<String>,
48}
49
50/// Configure code generation, then choose one explicit schema source.
51///
52/// Custom Rust syntax is checked before writing output. Generated structs preserve SQL names
53/// with Serde attributes. Additional derives and attributes apply to data types, not markers.
54#[derive(Debug, Clone)]
55pub struct Generator {
56    config: Config,
57}
58
59impl Default for Generator {
60    fn default() -> Self {
61        Self::new()
62    }
63}
64
65#[expect(
66    clippy::impl_trait_in_params,
67    reason = "Consuming configuration methods accept owned strings and borrowed literals without named generic parameters."
68)]
69#[expect(
70    clippy::wrong_self_convention,
71    reason = "Each from_* method consumes configured generation options and selects an input adapter."
72)]
73impl Generator {
74    #[must_use]
75    pub fn new() -> Self {
76        Self {
77            config: Config {
78                schemas: vec!["public".to_owned()],
79                prelude: String::new(),
80                derives: Vec::new(),
81                attributes: Vec::new(),
82                type_overrides: BTreeMap::new(),
83                type_attributes: BTreeMap::new(),
84                runtime_path: "::rp_supabase_client::schema".to_owned(),
85                not_null: BTreeMap::new(),
86                column_types: BTreeMap::new(),
87                json_types: BTreeMap::new(),
88                reexport_macros: false,
89                relationship_aliases: BTreeMap::new(),
90                strict_args: false,
91                strict_functions: BTreeSet::new(),
92            },
93        }
94    }
95
96    /// Replace the default `public` selection with these schemas.
97    #[must_use]
98    pub fn schemas(mut self, schemas: impl IntoIterator<Item = impl Into<String>>) -> Self {
99        self.config.schemas = schemas.into_iter().map(Into::into).collect();
100        self
101    }
102
103    /// Select a single schema instead of the default `public`.
104    #[must_use]
105    pub fn schema(self, schema: impl Into<String>) -> Self {
106        self.schemas([schema.into()])
107    }
108
109    /// Add Rust items or imports before the generated schema modules.
110    #[must_use]
111    pub fn prelude(mut self, prelude: impl Into<String>) -> Self {
112        self.config.prelude = prelude.into();
113        self
114    }
115
116    /// Add a derive path, such as `::typed_builder::TypedBuilder`.
117    #[must_use]
118    pub fn derive(mut self, derive: impl Into<String>) -> Self {
119        self.config.derives.push(derive.into());
120        self
121    }
122
123    /// Add an outer attribute, including its `#[...]` delimiters.
124    #[must_use]
125    pub fn attribute(mut self, attribute: impl Into<String>) -> Self {
126        self.config.attributes.push(attribute.into());
127        self
128    }
129
130    /// Add an attribute to one generated data type.
131    ///
132    /// Paths use generated Rust names, for example `public.tables.messages.Insert`.
133    /// Unknown paths fail generation instead of silently ignoring customization.
134    #[must_use]
135    pub fn type_attribute(
136        mut self,
137        generated_type: impl Into<String>,
138        attribute: impl Into<String>,
139    ) -> Self {
140        self.config
141            .type_attributes
142            .entry(generated_type.into())
143            .or_default()
144            .push(attribute.into());
145        self
146    }
147
148    /// Map a qualified SQL type to a Rust type that matches the `PostgREST` JSON representation.
149    #[must_use]
150    pub fn type_override(
151        mut self,
152        postgres_type: impl Into<String>,
153        rust_type: impl Into<String>,
154    ) -> Self {
155        self.config
156            .type_overrides
157            .insert(postgres_type.into(), rust_type.into());
158        self
159    }
160
161    /// Set the runtime module path when the client dependency has a Cargo alias.
162    #[must_use]
163    pub fn runtime_path(mut self, path: impl Into<String>) -> Self {
164        self.config.runtime_path = path.into();
165        self
166    }
167
168    /// Assert non-null fields on a view Row, composite, or RPC Record.
169    /// Unknown targets and SQL field names fail generation.
170    #[must_use]
171    pub fn not_null(
172        mut self,
173        target: impl Into<String>,
174        fields: impl IntoIterator<Item = impl Into<String>>,
175    ) -> Self {
176        self.config
177            .not_null
178            .insert(target.into(), fields.into_iter().map(Into::into).collect());
179        self
180    }
181
182    /// Override one SQL column's base type, preserving nullability and omission.
183    /// Accepts `public.table.column` or `public.tables.table.column`.
184    #[must_use]
185    pub fn column_type(mut self, target: impl Into<String>, rust_type: impl Into<String>) -> Self {
186        self.config
187            .column_types
188            .insert(target.into(), rust_type.into());
189        self
190    }
191
192    /// Decode a JSON field, RPC argument or return using a custom Serde type.
193    /// SQL arrays, nullability, and set-returning wrappers remain intact.
194    #[must_use]
195    pub fn json_type(mut self, target: impl Into<String>, rust_type: impl Into<String>) -> Self {
196        self.config
197            .json_types
198            .insert(target.into(), rust_type.into());
199        self
200    }
201
202    /// Export crate-local `select!` and `key!` wrappers with the runtime path filled.
203    /// Include the generated bindings at the schema crate's root.
204    #[must_use]
205    pub const fn reexport_macros(mut self) -> Self {
206        self.config.reexport_macros = true;
207        self
208    }
209
210    /// Add an alias for one canonical qualified relationship marker.
211    /// For example `public.tables.orders.relationships.orders_customer_fkey`.
212    #[must_use]
213    pub fn relationship_alias(
214        mut self,
215        target: impl Into<String>,
216        alias: impl Into<String>,
217    ) -> Self {
218        self.config
219            .relationship_aliases
220            .insert(target.into(), alias.into());
221        self
222    }
223
224    /// Require non-default RPC inputs unless explicitly marked nullable.
225    #[must_use]
226    pub const fn strict_args(mut self) -> Self {
227        self.config.strict_args = true;
228        self
229    }
230
231    /// Apply strict RPC input contracts to one generated function module.
232    #[must_use]
233    pub fn strict_args_for(mut self, target: impl Into<String>) -> Self {
234        self.config.strict_functions.insert(target.into());
235        self
236    }
237
238    /// Load a portable snapshot and register Cargo's file change detection.
239    ///
240    /// # Errors
241    /// Fails for unreadable, invalid, or unsupported-version snapshots.
242    #[expect(
243        clippy::print_stdout,
244        reason = "Cargo build scripts receive change instructions on stdout."
245    )]
246    pub fn from_snapshot(self, path: impl AsRef<Path>) -> Result<Bindings, Error> {
247        let path = path.as_ref();
248        cargo_path(path)?;
249        println!("cargo::rerun-if-changed={}", path.display());
250        let snapshot = parse_snapshot(&std::fs::read(path)?)?;
251        self.from_metadata(snapshot)
252    }
253
254    /// Generate from in-memory metadata. Does not emit Cargo change instructions.
255    ///
256    /// # Errors
257    /// Fails for unsupported versions, invalid metadata, or invalid custom Rust syntax.
258    pub fn from_metadata(self, mut snapshot: Snapshot) -> Result<Bindings, Error> {
259        validate_snapshot_version(snapshot.version)?;
260        if self.config.schemas.is_empty() || self.config.schemas.iter().any(String::is_empty) {
261            return Err(Error::Invalid(
262                "select at least one nonempty schema".to_owned(),
263            ));
264        }
265        for schema in &self.config.schemas {
266            if !snapshot.schemas.iter().any(|item| &item.name == schema) {
267                return Err(Error::Invalid(format!(
268                    "selected schema {schema:?} is absent from the snapshot"
269                )));
270            }
271        }
272        model::apply_not_null(&mut snapshot, &self.config.not_null)?;
273        let source = emitter::generate(&snapshot, &self.config)?;
274        Ok(Bindings { snapshot, source })
275    }
276
277    /// Acquire portable metadata without rendering Rust bindings.
278    ///
279    /// # Errors
280    /// Fails for empty schema selection, connection errors, or invalid catalog metadata.
281    #[cfg(feature = "database")]
282    pub fn snapshot_from_database(&self, url: &str) -> Result<Snapshot, Error> {
283        if self.config.schemas.is_empty() || self.config.schemas.iter().any(String::is_empty) {
284            return Err(Error::Invalid(
285                "select at least one nonempty schema".to_owned(),
286            ));
287        }
288        database::introspect(url, &self.config.schemas)
289    }
290
291    /// Introspect `PostgreSQL` directly. No CLI or Supabase service-role key is needed.
292    ///
293    /// Register migration paths in your build script, because Cargo cannot detect database DDL.
294    /// Remote connections should use `sslmode=require`; certificates and hostnames are verified.
295    ///
296    /// # Errors
297    /// Fails on connection, catalog, unsupported SQL type, or generation errors.
298    #[cfg(feature = "database")]
299    pub fn from_database(self, url: &str) -> Result<Bindings, Error> {
300        let snapshot = self.snapshot_from_database(url)?;
301        self.from_metadata(snapshot)
302    }
303
304    /// Read a database connection string from a host environment variable.
305    ///
306    /// Emits `rerun-if-env-changed` without printing the variable's secret value.
307    ///
308    /// # Errors
309    /// Fails if the variable is absent, invalid, or database introspection fails.
310    #[cfg(feature = "database")]
311    #[expect(
312        clippy::print_stdout,
313        reason = "Cargo build scripts receive change instructions on stdout."
314    )]
315    pub fn from_database_env(self, variable: &str) -> Result<Bindings, Error> {
316        if variable.is_empty() || variable.contains(['\n', '\r', '=']) {
317            return Err(Error::Invalid(
318                "invalid database environment variable name".to_owned(),
319            ));
320        }
321        println!("cargo::rerun-if-env-changed={variable}");
322        let url = std::env::var(variable).map_err(|_error| {
323            Error::Invalid(format!("set {variable} to a PostgreSQL connection string"))
324        })?;
325        self.from_database(&url)
326    }
327}
328
329/// Validated generated Rust source and the metadata used to produce it.
330#[derive(Debug, Clone)]
331pub struct Bindings {
332    snapshot: Snapshot,
333    source: String,
334}
335
336#[expect(
337    clippy::impl_trait_in_params,
338    reason = "Path arguments accept standard owned and borrowed path types."
339)]
340impl Bindings {
341    #[must_use]
342    pub fn source(&self) -> &str {
343        &self.source
344    }
345
346    #[must_use]
347    pub const fn snapshot(&self) -> &Snapshot {
348        &self.snapshot
349    }
350
351    /// Write source without changing the file when its content is identical.
352    ///
353    /// # Errors
354    /// Fails when reading or writing the output file fails.
355    pub fn write_to(&self, path: impl AsRef<Path>) -> Result<(), Error> {
356        write_if_changed(path.as_ref(), self.source.as_bytes())
357    }
358
359    /// Write source inside Cargo's `OUT_DIR`. The filename must be a single path component.
360    ///
361    /// # Errors
362    /// Fails outside a build script or if the filename or output path is invalid.
363    pub fn write_to_out_dir(&self, filename: &str) -> Result<PathBuf, Error> {
364        let path = Path::new(filename);
365        if path.components().count() != 1
366            || !matches!(
367                path.components().next(),
368                Some(std::path::Component::Normal(_))
369            )
370        {
371            return Err(Error::Invalid(
372                "OUT_DIR filename must be one normal path component".to_owned(),
373            ));
374        }
375        let output = std::env::var_os("OUT_DIR").ok_or_else(|| {
376            Error::Invalid("OUT_DIR is absent; call this method from build.rs".to_owned())
377        })?;
378        let output = PathBuf::from(output).join(filename);
379        self.write_to(&output)?;
380        Ok(output)
381    }
382
383    /// Export metadata for later offline builds. Call outside build.rs to update a committed snapshot.
384    ///
385    /// # Errors
386    /// Fails when encoding or writing the snapshot fails.
387    pub fn write_snapshot(&self, path: impl AsRef<Path>) -> Result<(), Error> {
388        self.snapshot.write_to(path)
389    }
390}
391
392fn validate_snapshot_version(version: u32) -> Result<(), Error> {
393    if version != SNAPSHOT_VERSION {
394        return Err(Error::Invalid(format!(
395            "snapshot version {version} is unsupported; expected {SNAPSHOT_VERSION}; \
396             regenerate the snapshot from PostgreSQL with the current supabase-codegen"
397        )));
398    }
399    Ok(())
400}
401
402fn parse_snapshot(bytes: &[u8]) -> Result<Snapshot, Error> {
403    #[derive(serde::Deserialize)]
404    struct VersionHeader {
405        version: u32,
406    }
407
408    let header: VersionHeader = serde_json::from_slice(bytes)?;
409    validate_snapshot_version(header.version)?;
410    Ok(serde_json::from_slice(bytes)?)
411}
412
413fn cargo_path(path: &Path) -> Result<(), Error> {
414    let Some(path) = path.to_str() else {
415        return Err(Error::Invalid("Cargo input paths must be UTF-8".to_owned()));
416    };
417    if path.contains(['\n', '\r']) {
418        return Err(Error::Invalid(
419            "Cargo input paths cannot contain line breaks".to_owned(),
420        ));
421    }
422    Ok(())
423}
424
425fn write_if_changed(path: &Path, contents: &[u8]) -> Result<(), Error> {
426    match std::fs::read(path) {
427        Ok(existing) if existing == contents => return Ok(()),
428        Ok(_) => {}
429        Err(error) if error.kind() == std::io::ErrorKind::NotFound => {}
430        Err(error) => return Err(error.into()),
431    }
432    std::fs::write(path, contents)?;
433    Ok(())
434}
435
436#[cfg(test)]
437#[expect(
438    clippy::expect_used,
439    reason = "Unexpected fixture decoding errors must print their diagnostic."
440)]
441mod metadata_tests {
442    use super::{Error, Generator, parse_snapshot};
443    use crate::model::{Snapshot, Table};
444
445    #[test]
446    fn legacy_snapshot_rejected_before_missing_relationship_metadata() {
447        assert!(matches!(
448            parse_snapshot(
449                br#"{"version":1,"schemas":[{"name":"public","enums":[],"composites":[],"functions":[],"tables":[{"name":"orders","kind":"table","columns":[]}]}]}"#,
450            ),
451            Err(Error::Invalid(_))
452        ));
453    }
454
455    #[test]
456    fn in_memory_metadata_rejects_legacy_version() {
457        assert!(matches!(
458            Generator::default().from_metadata(Snapshot {
459                version: 1,
460                schemas: Vec::new(),
461            }),
462            Err(Error::Invalid(_))
463        ));
464    }
465
466    #[test]
467    fn table_relationship_metadata_is_required_even_for_primary_key() {
468        let table = serde_json::json!({
469            "name": "orders",
470            "kind": "table",
471            "columns": [],
472            "primary_key": null,
473            "unique_keys": [],
474            "foreign_keys": [],
475            "is_partition": false
476        });
477        for field in ["primary_key", "unique_keys", "foreign_keys", "is_partition"] {
478            let mut incomplete = table.clone();
479            incomplete.as_object_mut().expect("object").remove(field);
480            assert!(matches!(
481                serde_json::from_value::<Table>(incomplete),
482                Err(error) if error.classify() == serde_json::error::Category::Data
483            ));
484        }
485        let decoded: Table = serde_json::from_value(table).expect("explicit empty facts are valid");
486        assert_eq!(decoded.primary_key, None);
487    }
488}