1mod grammar;
4
5use std::path::Path;
6
7use pest::Parser;
8use smol_str::SmolStr;
9use tracing::{debug, info};
10
11use crate::ast::*;
12use crate::error::{SchemaError, SchemaResult};
13
14pub use grammar::{PraxParser, Rule};
15
16use crate::ast::{
17 MssqlBlockOperation, Policy, PolicyCommand, PolicyType, Server, ServerGroup, ServerProperty,
18 ServerPropertyValue,
19};
20
21pub fn parse_schema(input: &str) -> SchemaResult<Schema> {
23 debug!(input_len = input.len(), "parse_schema() starting");
24 let pairs = PraxParser::parse(Rule::schema, input)
25 .map_err(|e| SchemaError::syntax(input.to_string(), 0, input.len(), e.to_string()))?;
26
27 let mut schema = Schema::new();
28 let mut current_doc: Option<Documentation> = None;
29
30 let schema_pair = pairs.into_iter().next().unwrap();
32
33 for pair in schema_pair.into_inner() {
34 match pair.as_rule() {
35 Rule::documentation => {
36 let span = pair.as_span();
37 let text = pair
38 .into_inner()
39 .map(|p| p.as_str().trim_start_matches("///").trim())
40 .collect::<Vec<_>>()
41 .join("\n");
42 current_doc = Some(Documentation::new(
43 text,
44 Span::new(span.start(), span.end()),
45 ));
46 }
47 Rule::model_def => {
48 let mut model = parse_model(pair)?;
49 if let Some(doc) = current_doc.take() {
50 model = model.with_documentation(doc);
51 }
52 schema.add_model(model);
53 }
54 Rule::enum_def => {
55 let mut e = parse_enum(pair)?;
56 if let Some(doc) = current_doc.take() {
57 e = e.with_documentation(doc);
58 }
59 schema.add_enum(e);
60 }
61 Rule::type_def => {
62 let mut t = parse_composite_type(pair)?;
63 if let Some(doc) = current_doc.take() {
64 t = t.with_documentation(doc);
65 }
66 schema.add_type(t);
67 }
68 Rule::view_def => {
69 let mut v = parse_view(pair)?;
70 if let Some(doc) = current_doc.take() {
71 v = v.with_documentation(doc);
72 }
73 schema.add_view(v);
74 }
75 Rule::raw_sql_def => {
76 let sql = parse_raw_sql(pair)?;
77 schema.add_raw_sql(sql);
78 }
79 Rule::server_group_def => {
80 let mut sg = parse_server_group(pair)?;
81 if let Some(doc) = current_doc.take() {
82 sg.set_documentation(doc);
83 }
84 schema.add_server_group(sg);
85 }
86 Rule::policy_def => {
87 let mut policy = parse_policy(pair)?;
88 if let Some(doc) = current_doc.take() {
89 policy = policy.with_documentation(doc);
90 }
91 schema.add_policy(policy);
92 }
93 Rule::datasource_def => {
94 let ds = parse_datasource(pair)?;
95 schema.set_datasource(ds);
96 current_doc = None;
97 }
98 Rule::generator_def => {
99 let generator = parse_generator(pair)?;
100 schema.add_generator(generator);
101 current_doc = None;
102 }
103 Rule::EOI => {}
104 _ => {}
105 }
106 }
107
108 info!(
109 models = schema.models.len(),
110 enums = schema.enums.len(),
111 types = schema.types.len(),
112 views = schema.views.len(),
113 generators = schema.generators.len(),
114 policies = schema.policies.len(),
115 "Schema parsed successfully"
116 );
117 Ok(schema)
118}
119
120pub fn parse_schema_file(path: impl AsRef<Path>) -> SchemaResult<Schema> {
122 let path = path.as_ref();
123 info!(path = %path.display(), "Loading schema file");
124 let content = std::fs::read_to_string(path).map_err(|e| SchemaError::IoError {
125 path: path.display().to_string(),
126 source: e,
127 })?;
128
129 parse_schema(&content)
130}
131
132fn parse_model(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<Model> {
134 let span = pair.as_span();
135 let mut inner = pair.into_inner();
136
137 let name_pair = inner.next().unwrap();
138 let name = Ident::new(
139 name_pair.as_str(),
140 Span::new(name_pair.as_span().start(), name_pair.as_span().end()),
141 );
142
143 let mut model = Model::new(name, Span::new(span.start(), span.end()));
144
145 for item in inner {
146 match item.as_rule() {
147 Rule::field_def => {
148 let field = parse_field(item)?;
149 model.add_field(field);
150 }
151 Rule::model_attribute => {
152 let attr = parse_attribute(item)?;
153 model.attributes.push(attr);
154 }
155 Rule::model_body_item => {
156 let inner_item = item.into_inner().next().unwrap();
158 match inner_item.as_rule() {
159 Rule::field_def => {
160 let field = parse_field(inner_item)?;
161 model.add_field(field);
162 }
163 Rule::model_attribute => {
164 let attr = parse_attribute(inner_item)?;
165 model.attributes.push(attr);
166 }
167 _ => {}
168 }
169 }
170 _ => {}
171 }
172 }
173
174 Ok(model)
175}
176
177fn parse_enum(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<Enum> {
179 let span = pair.as_span();
180 let mut inner = pair.into_inner();
181
182 let name_pair = inner.next().unwrap();
183 let name = Ident::new(
184 name_pair.as_str(),
185 Span::new(name_pair.as_span().start(), name_pair.as_span().end()),
186 );
187
188 let mut e = Enum::new(name, Span::new(span.start(), span.end()));
189
190 for item in inner {
191 match item.as_rule() {
192 Rule::enum_variant => {
193 let variant = parse_enum_variant(item)?;
194 e.add_variant(variant);
195 }
196 Rule::model_attribute => {
197 let attr = parse_attribute(item)?;
198 e.attributes.push(attr);
199 }
200 Rule::enum_body_item => {
201 let inner_item = item.into_inner().next().unwrap();
203 match inner_item.as_rule() {
204 Rule::enum_variant => {
205 let variant = parse_enum_variant(inner_item)?;
206 e.add_variant(variant);
207 }
208 Rule::model_attribute => {
209 let attr = parse_attribute(inner_item)?;
210 e.attributes.push(attr);
211 }
212 _ => {}
213 }
214 }
215 _ => {}
216 }
217 }
218
219 Ok(e)
220}
221
222fn parse_enum_variant(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<EnumVariant> {
224 let span = pair.as_span();
225 let mut inner = pair.into_inner();
226
227 let name_pair = inner.next().unwrap();
228 let name = Ident::new(
229 name_pair.as_str(),
230 Span::new(name_pair.as_span().start(), name_pair.as_span().end()),
231 );
232
233 let mut variant = EnumVariant::new(name, Span::new(span.start(), span.end()));
234
235 for item in inner {
236 if item.as_rule() == Rule::field_attribute {
237 let attr = parse_attribute(item)?;
238 variant.attributes.push(attr);
239 }
240 }
241
242 Ok(variant)
243}
244
245fn parse_composite_type(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<CompositeType> {
247 let span = pair.as_span();
248 let mut inner = pair.into_inner();
249
250 let name_pair = inner.next().unwrap();
251 let name = Ident::new(
252 name_pair.as_str(),
253 Span::new(name_pair.as_span().start(), name_pair.as_span().end()),
254 );
255
256 let mut t = CompositeType::new(name, Span::new(span.start(), span.end()));
257
258 for item in inner {
259 if item.as_rule() == Rule::field_def {
260 let field = parse_field(item)?;
261 t.add_field(field);
262 }
263 }
264
265 Ok(t)
266}
267
268fn parse_view(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<View> {
270 let span = pair.as_span();
271 let mut inner = pair.into_inner();
272
273 let name_pair = inner.next().unwrap();
274 let name = Ident::new(
275 name_pair.as_str(),
276 Span::new(name_pair.as_span().start(), name_pair.as_span().end()),
277 );
278
279 let mut v = View::new(name, Span::new(span.start(), span.end()));
280
281 for item in inner {
282 match item.as_rule() {
283 Rule::field_def => {
284 let field = parse_field(item)?;
285 v.add_field(field);
286 }
287 Rule::model_attribute => {
288 let attr = parse_attribute(item)?;
289 v.attributes.push(attr);
290 }
291 Rule::model_body_item => {
292 let inner_item = item.into_inner().next().unwrap();
294 match inner_item.as_rule() {
295 Rule::field_def => {
296 let field = parse_field(inner_item)?;
297 v.add_field(field);
298 }
299 Rule::model_attribute => {
300 let attr = parse_attribute(inner_item)?;
301 v.attributes.push(attr);
302 }
303 _ => {}
304 }
305 }
306 _ => {}
307 }
308 }
309
310 Ok(v)
311}
312
313fn parse_field(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<Field> {
315 let span = pair.as_span();
316 let mut inner = pair.into_inner();
317
318 let name_pair = inner.next().unwrap();
319 let name = Ident::new(
320 name_pair.as_str(),
321 Span::new(name_pair.as_span().start(), name_pair.as_span().end()),
322 );
323
324 let type_pair = inner.next().unwrap();
325 let (field_type, modifier) = parse_field_type(type_pair)?;
326
327 let mut attributes = vec![];
328 for item in inner {
329 if item.as_rule() == Rule::field_attribute {
330 let attr = parse_attribute(item)?;
331 attributes.push(attr);
332 }
333 }
334
335 Ok(Field::new(
336 name,
337 field_type,
338 modifier,
339 attributes,
340 Span::new(span.start(), span.end()),
341 ))
342}
343
344fn parse_field_type(
346 pair: pest::iterators::Pair<'_, Rule>,
347) -> SchemaResult<(FieldType, TypeModifier)> {
348 let mut type_name = String::new();
349 let mut modifier = TypeModifier::Required;
350
351 for item in pair.into_inner() {
352 match item.as_rule() {
353 Rule::type_name => {
354 type_name = item.as_str().to_string();
355 }
356 Rule::optional_marker => {
357 modifier = if modifier == TypeModifier::List {
358 TypeModifier::OptionalList
359 } else {
360 TypeModifier::Optional
361 };
362 }
363 Rule::list_marker => {
364 modifier = if modifier == TypeModifier::Optional {
365 TypeModifier::OptionalList
366 } else {
367 TypeModifier::List
368 };
369 }
370 _ => {}
371 }
372 }
373
374 let field_type = if let Some(scalar) = ScalarType::from_str(&type_name) {
375 FieldType::Scalar(scalar)
376 } else {
377 FieldType::Model(SmolStr::new(&type_name))
382 };
383
384 Ok((field_type, modifier))
385}
386
387fn parse_attribute(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<Attribute> {
389 let span = pair.as_span();
390 let mut inner = pair.into_inner();
391
392 let name_pair = inner.next().unwrap();
393 let name = Ident::new(
394 name_pair.as_str(),
395 Span::new(name_pair.as_span().start(), name_pair.as_span().end()),
396 );
397
398 let mut args = vec![];
399 for item in inner {
400 if item.as_rule() == Rule::attribute_args {
401 args = parse_attribute_args(item)?;
402 }
403 }
404
405 Ok(Attribute::new(
406 name,
407 args,
408 Span::new(span.start(), span.end()),
409 ))
410}
411
412fn parse_attribute_args(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<Vec<AttributeArg>> {
414 let mut args = vec![];
415
416 for item in pair.into_inner() {
417 if item.as_rule() == Rule::attribute_arg {
418 let arg = parse_attribute_arg(item)?;
419 args.push(arg);
420 }
421 }
422
423 Ok(args)
424}
425
426fn parse_attribute_arg(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<AttributeArg> {
428 let span = pair.as_span();
429 let mut inner = pair.into_inner();
430
431 let first = inner.next().unwrap();
432
433 if let Some(second) = inner.next() {
435 let name = Ident::new(
437 first.as_str(),
438 Span::new(first.as_span().start(), first.as_span().end()),
439 );
440 let value = parse_attribute_value(second)?;
441 Ok(AttributeArg::named(
442 name,
443 value,
444 Span::new(span.start(), span.end()),
445 ))
446 } else {
447 let value = parse_attribute_value(first)?;
449 Ok(AttributeArg::positional(
450 value,
451 Span::new(span.start(), span.end()),
452 ))
453 }
454}
455
456fn parse_attribute_value(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<AttributeValue> {
458 match pair.as_rule() {
459 Rule::string_literal => {
460 let s = pair.as_str();
461 let unquoted = &s[1..s.len() - 1];
463 Ok(AttributeValue::String(unquoted.to_string()))
464 }
465 Rule::number_literal => {
466 let s = pair.as_str();
467 if s.contains('.') {
468 Ok(AttributeValue::Float(s.parse().unwrap()))
469 } else {
470 Ok(AttributeValue::Int(s.parse().unwrap()))
471 }
472 }
473 Rule::boolean_literal => Ok(AttributeValue::Boolean(pair.as_str() == "true")),
474 Rule::identifier => Ok(AttributeValue::Ident(SmolStr::new(pair.as_str()))),
475 Rule::dotted_identifier => {
476 Ok(AttributeValue::String(pair.as_str().to_string()))
478 }
479 Rule::function_call => {
480 let mut inner = pair.into_inner();
481 let name = SmolStr::new(inner.next().unwrap().as_str());
482 let mut args = vec![];
483 for item in inner {
484 args.push(parse_attribute_value(item)?);
485 }
486 Ok(AttributeValue::Function(name, args))
487 }
488 Rule::field_ref_list => {
489 let refs: Vec<SmolStr> = pair
490 .into_inner()
491 .map(|p| SmolStr::new(p.as_str()))
492 .collect();
493 Ok(AttributeValue::FieldRefList(refs))
494 }
495 Rule::array_literal => {
496 let values: Result<Vec<_>, _> = pair.into_inner().map(parse_attribute_value).collect();
497 Ok(AttributeValue::Array(values?))
498 }
499 Rule::attribute_value => {
500 parse_attribute_value(pair.into_inner().next().unwrap())
502 }
503 _ => {
504 Ok(AttributeValue::Ident(SmolStr::new(pair.as_str())))
506 }
507 }
508}
509
510fn parse_raw_sql(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<RawSql> {
512 let mut inner = pair.into_inner();
513
514 let name = inner.next().unwrap().as_str();
515 let sql = inner.next().unwrap().as_str();
516
517 let name = name.trim().trim_matches('"');
520
521 let sql_content = sql
523 .trim_start_matches("\"\"\"")
524 .trim_end_matches("\"\"\"")
525 .trim();
526
527 Ok(RawSql::new(name, sql_content))
528}
529
530fn parse_server_group(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<ServerGroup> {
532 let span = pair.as_span();
533 let mut inner = pair.into_inner();
534
535 let name_pair = inner.next().unwrap();
536 let name = Ident::new(
537 name_pair.as_str(),
538 Span::new(name_pair.as_span().start(), name_pair.as_span().end()),
539 );
540
541 let mut server_group = ServerGroup::new(name, Span::new(span.start(), span.end()));
542
543 for item in inner {
544 match item.as_rule() {
545 Rule::server_group_item => {
546 let inner_item = item.into_inner().next().unwrap();
548 match inner_item.as_rule() {
549 Rule::server_def => {
550 let server = parse_server(inner_item)?;
551 server_group.add_server(server);
552 }
553 Rule::model_attribute => {
554 let attr = parse_attribute(inner_item)?;
555 server_group.add_attribute(attr);
556 }
557 _ => {}
558 }
559 }
560 Rule::server_def => {
561 let server = parse_server(item)?;
562 server_group.add_server(server);
563 }
564 Rule::model_attribute => {
565 let attr = parse_attribute(item)?;
566 server_group.add_attribute(attr);
567 }
568 _ => {}
569 }
570 }
571
572 Ok(server_group)
573}
574
575fn parse_server(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<Server> {
577 let span = pair.as_span();
578 let mut inner = pair.into_inner();
579
580 let name_pair = inner.next().unwrap();
581 let name = Ident::new(
582 name_pair.as_str(),
583 Span::new(name_pair.as_span().start(), name_pair.as_span().end()),
584 );
585
586 let mut server = Server::new(name, Span::new(span.start(), span.end()));
587
588 for item in inner {
589 if item.as_rule() == Rule::server_property {
590 let prop = parse_server_property(item)?;
591 server.add_property(prop);
592 }
593 }
594
595 Ok(server)
596}
597
598fn parse_server_property(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<ServerProperty> {
600 let span = pair.as_span();
601 let mut inner = pair.into_inner();
602
603 let key_pair = inner.next().unwrap();
604 let key = key_pair.as_str();
605
606 let value_pair = inner.next().unwrap();
607 let value = parse_server_property_value(value_pair)?;
608
609 Ok(ServerProperty::new(
610 key,
611 value,
612 Span::new(span.start(), span.end()),
613 ))
614}
615
616fn parse_generator(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<Generator> {
618 let span = pair.as_span();
619 let mut inner = pair.into_inner();
620
621 let name = inner.next().unwrap().as_str();
622 let mut generator = Generator::new(name, Span::new(span.start(), span.end()));
623
624 for prop in inner {
625 if prop.as_rule() == Rule::datasource_property {
626 let mut prop_inner = prop.into_inner();
627 let key = prop_inner.next().unwrap().as_str();
628 let value_pair = prop_inner.next().unwrap();
629
630 match key {
631 "provider" => {
632 let s = extract_datasource_string(&value_pair);
633 generator.provider = Some(SmolStr::new(s));
634 }
635 "output" => {
636 let s = extract_datasource_string(&value_pair);
637 generator.output = Some(SmolStr::new(s));
638 }
639 "generate" => {
640 generator.generate = parse_generator_toggle(&value_pair);
641 }
642 _ => {
643 let val = parse_generator_value(&value_pair);
644 generator.properties.insert(SmolStr::new(key), val);
645 }
646 }
647 }
648 }
649
650 Ok(generator)
651}
652
653fn parse_generator_toggle(pair: &pest::iterators::Pair<'_, Rule>) -> GeneratorToggle {
655 match pair.as_rule() {
656 Rule::env_function => {
657 let env_var = pair
658 .clone()
659 .into_inner()
660 .next()
661 .map(|p| {
662 let s = p.as_str();
663 SmolStr::new(&s[1..s.len() - 1])
664 })
665 .unwrap_or_default();
666 GeneratorToggle::Env(env_var)
667 }
668 Rule::datasource_value => {
669 let inner = pair.clone().into_inner().next().unwrap();
670 parse_generator_toggle(&inner)
671 }
672 _ => {
673 let s = pair.as_str().trim().trim_matches('"');
674 match s {
675 "true" => GeneratorToggle::Literal(true),
676 "false" => GeneratorToggle::Literal(false),
677 _ => GeneratorToggle::Literal(false),
678 }
679 }
680 }
681}
682
683fn parse_generator_value(pair: &pest::iterators::Pair<'_, Rule>) -> GeneratorValue {
685 match pair.as_rule() {
686 Rule::env_function => {
687 let env_var = pair
688 .clone()
689 .into_inner()
690 .next()
691 .map(|p| {
692 let s = p.as_str();
693 SmolStr::new(&s[1..s.len() - 1])
694 })
695 .unwrap_or_default();
696 GeneratorValue::Env(env_var)
697 }
698 Rule::datasource_value => {
699 let inner = pair.clone().into_inner().next().unwrap();
700 parse_generator_value(&inner)
701 }
702 Rule::string_literal => {
703 let s = pair.as_str();
704 GeneratorValue::String(SmolStr::new(&s[1..s.len() - 1]))
705 }
706 _ => {
707 let s = pair.as_str().trim().trim_matches('"');
708 match s {
709 "true" => GeneratorValue::Bool(true),
710 "false" => GeneratorValue::Bool(false),
711 _ => GeneratorValue::Ident(SmolStr::new(s)),
712 }
713 }
714 }
715}
716
717fn parse_datasource(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<Datasource> {
719 let span = pair.as_span();
720 let mut inner = pair.into_inner();
721
722 let name_pair = inner.next().unwrap();
723 let name = name_pair.as_str();
724
725 let mut datasource = Datasource::new(
726 name,
727 DatabaseProvider::PostgreSQL,
728 Span::new(span.start(), span.end()),
729 );
730
731 for prop in inner {
732 if prop.as_rule() == Rule::datasource_property {
733 let mut prop_inner = prop.into_inner();
734 let key = prop_inner.next().unwrap().as_str();
735 let value_pair = prop_inner.next().unwrap();
736
737 match key {
738 "provider" => {
739 let provider_str = extract_datasource_string(&value_pair);
740 if let Some(provider) = DatabaseProvider::from_str(&provider_str) {
741 datasource.provider = provider;
742 }
743 }
744 "url" => {
745 match value_pair.as_rule() {
746 Rule::env_function => {
747 let env_var = value_pair
749 .into_inner()
750 .next()
751 .map(|p| {
752 let s = p.as_str();
753 s[1..s.len() - 1].to_string()
754 })
755 .unwrap_or_default();
756 datasource.url_env = Some(SmolStr::new(env_var));
757 }
758 Rule::string_literal => {
759 let s = value_pair.as_str();
760 let url = &s[1..s.len() - 1];
761 datasource.url = Some(SmolStr::new(url));
762 }
763 _ => {}
764 }
765 }
766 "extensions" => {
767 if value_pair.as_rule() == Rule::extension_array {
768 for ext_item in value_pair.into_inner() {
769 if ext_item.as_rule() == Rule::extension_item {
770 let ext = parse_extension_item(
771 ext_item,
772 Span::new(span.start(), span.end()),
773 )?;
774 datasource.add_extension(ext);
775 }
776 }
777 }
778 }
779 _ => {
780 let value_str = extract_datasource_string(&value_pair);
782 datasource.add_property(key, value_str);
783 }
784 }
785 }
786 }
787
788 Ok(datasource)
789}
790
791fn parse_extension_item(
793 pair: pest::iterators::Pair<'_, Rule>,
794 span: Span,
795) -> SchemaResult<PostgresExtension> {
796 let mut inner = pair.into_inner();
797 let name = inner.next().unwrap().as_str();
798 let mut ext = PostgresExtension::new(name, span);
799
800 if let Some(args_pair) = inner.next()
802 && args_pair.as_rule() == Rule::extension_args
803 {
804 for arg in args_pair.into_inner() {
805 if arg.as_rule() == Rule::extension_arg {
806 let mut arg_inner = arg.into_inner();
807 let arg_key = arg_inner.next().unwrap().as_str();
808 let arg_value_pair = arg_inner.next().unwrap();
809 let arg_value = {
810 let s = arg_value_pair.as_str();
811 &s[1..s.len() - 1]
812 };
813
814 match arg_key {
815 "schema" => {
816 ext = ext.with_schema(arg_value);
817 }
818 "version" => {
819 ext = ext.with_version(arg_value);
820 }
821 _ => {}
822 }
823 }
824 }
825 }
826
827 Ok(ext)
828}
829
830fn extract_datasource_string(pair: &pest::iterators::Pair<'_, Rule>) -> String {
832 match pair.as_rule() {
833 Rule::string_literal => {
834 let s = pair.as_str();
835 s[1..s.len() - 1].to_string()
836 }
837 Rule::identifier => pair.as_str().to_string(),
838 Rule::datasource_value => {
839 if let Some(inner) = pair.clone().into_inner().next() {
840 extract_datasource_string(&inner)
841 } else {
842 pair.as_str().to_string()
843 }
844 }
845 _ => pair.as_str().to_string(),
846 }
847}
848
849fn extract_string_from_arg(pair: pest::iterators::Pair<'_, Rule>) -> String {
851 match pair.as_rule() {
852 Rule::string_literal => {
853 let s = pair.as_str();
854 s[1..s.len() - 1].to_string()
855 }
856 Rule::attribute_value => {
857 if let Some(inner) = pair.into_inner().next() {
859 extract_string_from_arg(inner)
860 } else {
861 String::new()
862 }
863 }
864 _ => pair.as_str().to_string(),
865 }
866}
867
868fn parse_server_property_value(
870 pair: pest::iterators::Pair<'_, Rule>,
871) -> SchemaResult<ServerPropertyValue> {
872 match pair.as_rule() {
873 Rule::string_literal => {
874 let s = pair.as_str();
875 let unquoted = &s[1..s.len() - 1];
877 Ok(ServerPropertyValue::String(unquoted.to_string()))
878 }
879 Rule::number_literal => {
880 let s = pair.as_str();
881 Ok(ServerPropertyValue::Number(s.parse().unwrap_or(0.0)))
882 }
883 Rule::boolean_literal => Ok(ServerPropertyValue::Boolean(pair.as_str() == "true")),
884 Rule::identifier => Ok(ServerPropertyValue::Identifier(pair.as_str().to_string())),
885 Rule::function_call => {
886 let mut inner = pair.into_inner();
888 let func_name = inner.next().unwrap().as_str();
889 if func_name == "env"
890 && let Some(arg) = inner.next()
891 {
892 let var_name = extract_string_from_arg(arg);
893 return Ok(ServerPropertyValue::EnvVar(var_name));
894 }
895 Ok(ServerPropertyValue::Identifier(func_name.to_string()))
897 }
898 Rule::array_literal => {
899 let values: Result<Vec<_>, _> =
900 pair.into_inner().map(parse_server_property_value).collect();
901 Ok(ServerPropertyValue::Array(values?))
902 }
903 Rule::attribute_value => {
904 parse_server_property_value(pair.into_inner().next().unwrap())
906 }
907 _ => {
908 Ok(ServerPropertyValue::Identifier(pair.as_str().to_string()))
910 }
911 }
912}
913
914fn parse_policy(pair: pest::iterators::Pair<'_, Rule>) -> SchemaResult<Policy> {
916 let span = pair.as_span();
917 let mut inner = pair.into_inner();
918
919 let name_pair = inner.next().unwrap();
921 let name = Ident::new(
922 name_pair.as_str(),
923 Span::new(name_pair.as_span().start(), name_pair.as_span().end()),
924 );
925
926 let table_pair = inner.next().unwrap();
928 let table = Ident::new(
929 table_pair.as_str(),
930 Span::new(table_pair.as_span().start(), table_pair.as_span().end()),
931 );
932
933 let mut policy = Policy::new(name, table, Span::new(span.start(), span.end()));
934 policy.commands = vec![];
936
937 for item in inner {
938 match item.as_rule() {
939 Rule::policy_item => {
940 let inner_item = item.into_inner().next().unwrap();
941 parse_policy_item(&mut policy, inner_item)?;
942 }
943 Rule::policy_for
944 | Rule::policy_to
945 | Rule::policy_as
946 | Rule::policy_using
947 | Rule::policy_check => {
948 parse_policy_item(&mut policy, item)?;
949 }
950 _ => {}
951 }
952 }
953
954 if policy.commands.is_empty() {
956 policy.commands.push(PolicyCommand::All);
957 }
958
959 Ok(policy)
960}
961
962fn parse_policy_item(
964 policy: &mut Policy,
965 pair: pest::iterators::Pair<'_, Rule>,
966) -> SchemaResult<()> {
967 match pair.as_rule() {
968 Rule::policy_for => {
969 let inner = pair.into_inner().next().unwrap();
970 match inner.as_rule() {
971 Rule::policy_command => {
972 if let Some(cmd) = PolicyCommand::from_str(inner.as_str()) {
973 policy.add_command(cmd);
974 }
975 }
976 Rule::policy_command_list => {
977 for cmd_pair in inner.into_inner() {
978 if cmd_pair.as_rule() == Rule::policy_command
979 && let Some(cmd) = PolicyCommand::from_str(cmd_pair.as_str())
980 {
981 policy.add_command(cmd);
982 }
983 }
984 }
985 _ => {}
986 }
987 }
988 Rule::policy_to => {
989 let inner = pair.into_inner().next().unwrap();
990 match inner.as_rule() {
991 Rule::identifier => {
992 policy.add_role(inner.as_str());
993 }
994 Rule::policy_role_list => {
995 for role_pair in inner.into_inner() {
996 if role_pair.as_rule() == Rule::identifier {
997 policy.add_role(role_pair.as_str());
998 }
999 }
1000 }
1001 _ => {}
1002 }
1003 }
1004 Rule::policy_as => {
1005 let inner = pair.into_inner().next().unwrap();
1006 if inner.as_rule() == Rule::policy_type
1007 && let Some(policy_type) = PolicyType::from_str(inner.as_str())
1008 {
1009 policy.policy_type = policy_type;
1010 }
1011 }
1012 Rule::policy_using => {
1013 let inner = pair.into_inner().next().unwrap();
1014 let expr = extract_policy_expression(&inner);
1015 policy.using_expr = Some(expr);
1016 }
1017 Rule::policy_check => {
1018 let inner = pair.into_inner().next().unwrap();
1019 let expr = extract_policy_expression(&inner);
1020 policy.check_expr = Some(expr);
1021 }
1022 Rule::policy_mssql_schema => {
1023 let inner = pair.into_inner().next().unwrap();
1024 if inner.as_rule() == Rule::string_literal {
1025 let s = inner.as_str();
1026 let schema = &s[1..s.len() - 1]; policy.mssql_schema = Some(SmolStr::new(schema));
1028 }
1029 }
1030 Rule::policy_mssql_block => {
1031 let inner = pair.into_inner().next().unwrap();
1032 match inner.as_rule() {
1033 Rule::mssql_block_op => {
1034 if let Some(op) = MssqlBlockOperation::from_str(inner.as_str()) {
1035 policy.add_mssql_block_operation(op);
1036 }
1037 }
1038 Rule::mssql_block_op_list => {
1039 for op_pair in inner.into_inner() {
1040 if op_pair.as_rule() == Rule::mssql_block_op
1041 && let Some(op) = MssqlBlockOperation::from_str(op_pair.as_str())
1042 {
1043 policy.add_mssql_block_operation(op);
1044 }
1045 }
1046 }
1047 _ => {}
1048 }
1049 }
1050 _ => {}
1051 }
1052 Ok(())
1053}
1054
1055fn extract_policy_expression(pair: &pest::iterators::Pair<'_, Rule>) -> String {
1057 let s = pair.as_str();
1058 match pair.as_rule() {
1059 Rule::multiline_string => {
1060 s.trim_start_matches("\"\"\"")
1062 .trim_end_matches("\"\"\"")
1063 .trim()
1064 .to_string()
1065 }
1066 Rule::string_literal => {
1067 s[1..s.len() - 1].to_string()
1069 }
1070 _ => s.to_string(),
1071 }
1072}
1073
1074#[cfg(test)]
1075mod tests {
1076 use super::*;
1077
1078 #[test]
1081 fn test_parse_simple_model() {
1082 let schema = parse_schema(
1083 r#"
1084 model User {
1085 id Int @id @auto
1086 email String @unique
1087 name String?
1088 }
1089 "#,
1090 )
1091 .unwrap();
1092
1093 assert_eq!(schema.models.len(), 1);
1094 let user = schema.get_model("User").unwrap();
1095 assert_eq!(user.fields.len(), 3);
1096 assert!(user.get_field("id").unwrap().is_id());
1097 assert!(user.get_field("email").unwrap().is_unique());
1098 assert!(user.get_field("name").unwrap().is_optional());
1099 }
1100
1101 #[test]
1102 fn test_parse_model_name() {
1103 let schema = parse_schema(
1104 r#"
1105 model BlogPost {
1106 id Int @id
1107 }
1108 "#,
1109 )
1110 .unwrap();
1111
1112 assert!(schema.get_model("BlogPost").is_some());
1113 }
1114
1115 #[test]
1116 fn test_parse_multiple_models() {
1117 let schema = parse_schema(
1118 r#"
1119 model User {
1120 id Int @id
1121 }
1122
1123 model Post {
1124 id Int @id
1125 }
1126
1127 model Comment {
1128 id Int @id
1129 }
1130 "#,
1131 )
1132 .unwrap();
1133
1134 assert_eq!(schema.models.len(), 3);
1135 assert!(schema.get_model("User").is_some());
1136 assert!(schema.get_model("Post").is_some());
1137 assert!(schema.get_model("Comment").is_some());
1138 }
1139
1140 #[test]
1143 fn test_parse_all_scalar_types() {
1144 let schema = parse_schema(
1145 r#"
1146 model AllTypes {
1147 id Int @id
1148 big BigInt
1149 float_f Float
1150 decimal Decimal
1151 str String
1152 bool Boolean
1153 datetime DateTime
1154 date Date
1155 time Time
1156 json Json
1157 bytes Bytes
1158 uuid Uuid
1159 cuid Cuid
1160 cuid2 Cuid2
1161 nanoid NanoId
1162 ulid Ulid
1163 }
1164 "#,
1165 )
1166 .unwrap();
1167
1168 let model = schema.get_model("AllTypes").unwrap();
1169 assert_eq!(model.fields.len(), 16);
1170
1171 assert!(matches!(
1172 model.get_field("id").unwrap().field_type,
1173 FieldType::Scalar(ScalarType::Int)
1174 ));
1175 assert!(matches!(
1176 model.get_field("big").unwrap().field_type,
1177 FieldType::Scalar(ScalarType::BigInt)
1178 ));
1179 assert!(matches!(
1180 model.get_field("str").unwrap().field_type,
1181 FieldType::Scalar(ScalarType::String)
1182 ));
1183 assert!(matches!(
1184 model.get_field("bool").unwrap().field_type,
1185 FieldType::Scalar(ScalarType::Boolean)
1186 ));
1187 assert!(matches!(
1188 model.get_field("datetime").unwrap().field_type,
1189 FieldType::Scalar(ScalarType::DateTime)
1190 ));
1191 assert!(matches!(
1192 model.get_field("uuid").unwrap().field_type,
1193 FieldType::Scalar(ScalarType::Uuid)
1194 ));
1195 assert!(matches!(
1196 model.get_field("cuid").unwrap().field_type,
1197 FieldType::Scalar(ScalarType::Cuid)
1198 ));
1199 assert!(matches!(
1200 model.get_field("cuid2").unwrap().field_type,
1201 FieldType::Scalar(ScalarType::Cuid2)
1202 ));
1203 assert!(matches!(
1204 model.get_field("nanoid").unwrap().field_type,
1205 FieldType::Scalar(ScalarType::NanoId)
1206 ));
1207 assert!(matches!(
1208 model.get_field("ulid").unwrap().field_type,
1209 FieldType::Scalar(ScalarType::Ulid)
1210 ));
1211 }
1212
1213 #[test]
1214 fn test_parse_optional_field() {
1215 let schema = parse_schema(
1216 r#"
1217 model User {
1218 id Int @id
1219 bio String?
1220 age Int?
1221 }
1222 "#,
1223 )
1224 .unwrap();
1225
1226 let user = schema.get_model("User").unwrap();
1227 assert!(!user.get_field("id").unwrap().is_optional());
1228 assert!(user.get_field("bio").unwrap().is_optional());
1229 assert!(user.get_field("age").unwrap().is_optional());
1230 }
1231
1232 #[test]
1233 fn test_parse_list_field() {
1234 let schema = parse_schema(
1235 r#"
1236 model User {
1237 id Int @id
1238 tags String[]
1239 posts Post[]
1240 }
1241 "#,
1242 )
1243 .unwrap();
1244
1245 let user = schema.get_model("User").unwrap();
1246 assert!(user.get_field("tags").unwrap().is_list());
1247 assert!(user.get_field("posts").unwrap().is_list());
1248 }
1249
1250 #[test]
1251 fn test_parse_optional_list_field() {
1252 let schema = parse_schema(
1253 r#"
1254 model User {
1255 id Int @id
1256 metadata String[]?
1257 }
1258 "#,
1259 )
1260 .unwrap();
1261
1262 let user = schema.get_model("User").unwrap();
1263 let metadata = user.get_field("metadata").unwrap();
1264 assert!(metadata.is_list());
1265 assert!(metadata.is_optional());
1266 }
1267
1268 #[test]
1271 fn test_parse_id_attribute() {
1272 let schema = parse_schema(
1273 r#"
1274 model User {
1275 id Int @id
1276 }
1277 "#,
1278 )
1279 .unwrap();
1280
1281 let user = schema.get_model("User").unwrap();
1282 assert!(user.get_field("id").unwrap().is_id());
1283 }
1284
1285 #[test]
1286 fn test_parse_unique_attribute() {
1287 let schema = parse_schema(
1288 r#"
1289 model User {
1290 id Int @id
1291 email String @unique
1292 }
1293 "#,
1294 )
1295 .unwrap();
1296
1297 let user = schema.get_model("User").unwrap();
1298 assert!(user.get_field("email").unwrap().is_unique());
1299 }
1300
1301 #[test]
1302 fn test_parse_default_int() {
1303 let schema = parse_schema(
1304 r#"
1305 model Counter {
1306 id Int @id
1307 count Int @default(0)
1308 }
1309 "#,
1310 )
1311 .unwrap();
1312
1313 let counter = schema.get_model("Counter").unwrap();
1314 let count_field = counter.get_field("count").unwrap();
1315 let attrs = count_field.extract_attributes();
1316 assert!(attrs.default.is_some());
1317 assert_eq!(attrs.default.unwrap().as_int(), Some(0));
1318 }
1319
1320 #[test]
1321 fn test_parse_default_string() {
1322 let schema = parse_schema(
1323 r#"
1324 model User {
1325 id Int @id
1326 status String @default("active")
1327 }
1328 "#,
1329 )
1330 .unwrap();
1331
1332 let user = schema.get_model("User").unwrap();
1333 let status = user.get_field("status").unwrap();
1334 let attrs = status.extract_attributes();
1335 assert!(attrs.default.is_some());
1336 assert_eq!(attrs.default.unwrap().as_string(), Some("active"));
1337 }
1338
1339 #[test]
1340 fn test_parse_default_boolean() {
1341 let schema = parse_schema(
1342 r#"
1343 model Post {
1344 id Int @id
1345 published Boolean @default(false)
1346 }
1347 "#,
1348 )
1349 .unwrap();
1350
1351 let post = schema.get_model("Post").unwrap();
1352 let published = post.get_field("published").unwrap();
1353 let attrs = published.extract_attributes();
1354 assert!(attrs.default.is_some());
1355 assert_eq!(attrs.default.unwrap().as_bool(), Some(false));
1356 }
1357
1358 #[test]
1359 fn test_parse_default_function() {
1360 let schema = parse_schema(
1361 r#"
1362 model User {
1363 id Int @id
1364 createdAt DateTime @default(now())
1365 }
1366 "#,
1367 )
1368 .unwrap();
1369
1370 let user = schema.get_model("User").unwrap();
1371 let created_at = user.get_field("createdAt").unwrap();
1372 let attrs = created_at.extract_attributes();
1373 assert!(attrs.default.is_some());
1374 if let Some(AttributeValue::Function(name, _)) = attrs.default {
1375 assert_eq!(name.as_str(), "now");
1376 } else {
1377 panic!("Expected function default");
1378 }
1379 }
1380
1381 #[test]
1382 fn test_parse_updated_at_attribute() {
1383 let schema = parse_schema(
1384 r#"
1385 model User {
1386 id Int @id
1387 updatedAt DateTime @updated_at
1388 }
1389 "#,
1390 )
1391 .unwrap();
1392
1393 let user = schema.get_model("User").unwrap();
1394 let updated_at = user.get_field("updatedAt").unwrap();
1395 let attrs = updated_at.extract_attributes();
1396 assert!(attrs.is_updated_at);
1397 }
1398
1399 #[test]
1400 fn test_parse_map_attribute() {
1401 let schema = parse_schema(
1402 r#"
1403 model User {
1404 id Int @id
1405 email String @map("email_address")
1406 }
1407 "#,
1408 )
1409 .unwrap();
1410
1411 let user = schema.get_model("User").unwrap();
1412 let email = user.get_field("email").unwrap();
1413 let attrs = email.extract_attributes();
1414 assert_eq!(attrs.map, Some("email_address".to_string()));
1415 }
1416
1417 #[test]
1418 fn test_parse_multiple_attributes() {
1419 let schema = parse_schema(
1420 r#"
1421 model User {
1422 id Int @id @auto
1423 email String @unique @index
1424 }
1425 "#,
1426 )
1427 .unwrap();
1428
1429 let user = schema.get_model("User").unwrap();
1430 let id = user.get_field("id").unwrap();
1431 let email = user.get_field("email").unwrap();
1432
1433 let id_attrs = id.extract_attributes();
1434 assert!(id_attrs.is_id);
1435 assert!(id_attrs.is_auto);
1436
1437 let email_attrs = email.extract_attributes();
1438 assert!(email_attrs.is_unique);
1439 assert!(email_attrs.is_indexed);
1440 }
1441
1442 #[test]
1445 fn test_parse_model_map_attribute() {
1446 let schema = parse_schema(
1447 r#"
1448 model User {
1449 id Int @id
1450
1451 @@map("app_users")
1452 }
1453 "#,
1454 )
1455 .unwrap();
1456
1457 let user = schema.get_model("User").unwrap();
1458 assert_eq!(user.table_name(), "app_users");
1459 }
1460
1461 #[test]
1462 fn test_parse_model_index_attribute() {
1463 let schema = parse_schema(
1464 r#"
1465 model User {
1466 id Int @id
1467 email String
1468 name String
1469
1470 @@index([email, name])
1471 }
1472 "#,
1473 )
1474 .unwrap();
1475
1476 let user = schema.get_model("User").unwrap();
1477 assert!(user.has_attribute("index"));
1478 }
1479
1480 #[test]
1481 fn test_parse_composite_primary_key() {
1482 let schema = parse_schema(
1483 r#"
1484 model PostTag {
1485 postId Int
1486 tagId Int
1487
1488 @@id([postId, tagId])
1489 }
1490 "#,
1491 )
1492 .unwrap();
1493
1494 let post_tag = schema.get_model("PostTag").unwrap();
1495 assert!(post_tag.has_attribute("id"));
1496 }
1497
1498 #[test]
1501 fn test_parse_enum() {
1502 let schema = parse_schema(
1503 r#"
1504 enum Role {
1505 User
1506 Admin
1507 Moderator
1508 }
1509 "#,
1510 )
1511 .unwrap();
1512
1513 assert_eq!(schema.enums.len(), 1);
1514 let role = schema.get_enum("Role").unwrap();
1515 assert_eq!(role.variants.len(), 3);
1516 }
1517
1518 #[test]
1519 fn test_parse_enum_variant_names() {
1520 let schema = parse_schema(
1521 r#"
1522 enum Status {
1523 Pending
1524 Active
1525 Completed
1526 Cancelled
1527 }
1528 "#,
1529 )
1530 .unwrap();
1531
1532 let status = schema.get_enum("Status").unwrap();
1533 assert!(status.get_variant("Pending").is_some());
1534 assert!(status.get_variant("Active").is_some());
1535 assert!(status.get_variant("Completed").is_some());
1536 assert!(status.get_variant("Cancelled").is_some());
1537 }
1538
1539 #[test]
1540 fn test_parse_enum_with_map() {
1541 let schema = parse_schema(
1542 r#"
1543 enum Role {
1544 User @map("USER")
1545 Admin @map("ADMINISTRATOR")
1546 }
1547 "#,
1548 )
1549 .unwrap();
1550
1551 let role = schema.get_enum("Role").unwrap();
1552 let user_variant = role.get_variant("User").unwrap();
1553 assert_eq!(user_variant.db_value(), "USER");
1554
1555 let admin_variant = role.get_variant("Admin").unwrap();
1556 assert_eq!(admin_variant.db_value(), "ADMINISTRATOR");
1557 }
1558
1559 #[test]
1562 fn test_parse_one_to_many_relation() {
1563 let schema = parse_schema(
1564 r#"
1565 model User {
1566 id Int @id
1567 posts Post[]
1568 }
1569
1570 model Post {
1571 id Int @id
1572 authorId Int
1573 author User @relation(fields: [authorId], references: [id])
1574 }
1575 "#,
1576 )
1577 .unwrap();
1578
1579 let user = schema.get_model("User").unwrap();
1580 let post = schema.get_model("Post").unwrap();
1581
1582 assert!(user.get_field("posts").unwrap().is_list());
1583 assert!(post.get_field("author").unwrap().is_relation());
1584 }
1585
1586 #[test]
1587 fn test_parse_relation_with_actions() {
1588 let schema = parse_schema(
1589 r#"
1590 model Post {
1591 id Int @id
1592 authorId Int
1593 author User @relation(fields: [authorId], references: [id], onDelete: Cascade, onUpdate: Restrict)
1594 }
1595
1596 model User {
1597 id Int @id
1598 posts Post[]
1599 }
1600 "#,
1601 )
1602 .unwrap();
1603
1604 let post = schema.get_model("Post").unwrap();
1605 let author = post.get_field("author").unwrap();
1606 let attrs = author.extract_attributes();
1607
1608 assert!(attrs.relation.is_some());
1609 let rel = attrs.relation.unwrap();
1610 assert_eq!(rel.on_delete, Some(ReferentialAction::Cascade));
1611 assert_eq!(rel.on_update, Some(ReferentialAction::Restrict));
1612 }
1613
1614 #[test]
1617 fn test_parse_model_documentation() {
1618 let schema = parse_schema(
1619 r#"/// Represents a user in the system
1620model User {
1621 id Int @id
1622}"#,
1623 )
1624 .unwrap();
1625
1626 let user = schema.get_model("User").unwrap();
1627 if let Some(doc) = &user.documentation {
1630 assert!(doc.text.contains("user"));
1631 }
1632 }
1633
1634 #[test]
1637 fn test_parse_complete_schema() {
1638 let schema = parse_schema(
1639 r#"
1640 /// User model
1641 model User {
1642 id Int @id @auto
1643 email String @unique
1644 name String?
1645 role Role @default(User)
1646 posts Post[]
1647 profile Profile?
1648 createdAt DateTime @default(now())
1649 updatedAt DateTime @updated_at
1650
1651 @@map("users")
1652 @@index([email])
1653 }
1654
1655 model Post {
1656 id Int @id @auto
1657 title String
1658 content String?
1659 published Boolean @default(false)
1660 authorId Int
1661 author User @relation(fields: [authorId], references: [id])
1662 tags Tag[]
1663 createdAt DateTime @default(now())
1664
1665 @@index([authorId])
1666 }
1667
1668 model Profile {
1669 id Int @id @auto
1670 bio String?
1671 userId Int @unique
1672 user User @relation(fields: [userId], references: [id])
1673 }
1674
1675 model Tag {
1676 id Int @id @auto
1677 name String @unique
1678 posts Post[]
1679 }
1680
1681 enum Role {
1682 User
1683 Admin
1684 Moderator
1685 }
1686 "#,
1687 )
1688 .unwrap();
1689
1690 assert_eq!(schema.models.len(), 4);
1692 assert!(schema.get_model("User").is_some());
1693 assert!(schema.get_model("Post").is_some());
1694 assert!(schema.get_model("Profile").is_some());
1695 assert!(schema.get_model("Tag").is_some());
1696
1697 assert_eq!(schema.enums.len(), 1);
1699 assert!(schema.get_enum("Role").is_some());
1700
1701 let user = schema.get_model("User").unwrap();
1703 assert_eq!(user.table_name(), "users");
1704 assert_eq!(user.fields.len(), 8);
1705 assert!(user.has_attribute("index"));
1706
1707 let post = schema.get_model("Post").unwrap();
1709 assert!(post.get_field("author").unwrap().is_relation());
1710 }
1711
1712 #[test]
1715 fn test_parse_invalid_syntax() {
1716 let result = parse_schema("model { broken }");
1717 assert!(result.is_err());
1718 }
1719
1720 #[test]
1721 fn test_parse_empty_schema() {
1722 let schema = parse_schema("").unwrap();
1723 assert!(schema.models.is_empty());
1724 assert!(schema.enums.is_empty());
1725 }
1726
1727 #[test]
1728 fn test_parse_whitespace_only() {
1729 let schema = parse_schema(" \n\t \n ").unwrap();
1730 assert!(schema.models.is_empty());
1731 }
1732
1733 #[test]
1734 fn test_parse_comments_only() {
1735 let schema = parse_schema(
1736 r#"
1737 // This is a comment
1738 // Another comment
1739 "#,
1740 )
1741 .unwrap();
1742 assert!(schema.models.is_empty());
1743 }
1744
1745 #[test]
1748 fn test_parse_model_with_no_fields() {
1749 let result = parse_schema(
1751 r#"
1752 model Empty {
1753 }
1754 "#,
1755 );
1756 let _ = result;
1758 }
1759
1760 #[test]
1761 fn test_parse_long_identifier() {
1762 let schema = parse_schema(
1763 r#"
1764 model VeryLongModelNameThatIsStillValid {
1765 someVeryLongFieldNameThatShouldWork Int @id
1766 }
1767 "#,
1768 )
1769 .unwrap();
1770
1771 assert!(
1772 schema
1773 .get_model("VeryLongModelNameThatIsStillValid")
1774 .is_some()
1775 );
1776 }
1777
1778 #[test]
1779 fn test_parse_underscore_identifiers() {
1780 let schema = parse_schema(
1781 r#"
1782 model user_account {
1783 user_id Int @id
1784 created_at DateTime
1785 }
1786 "#,
1787 )
1788 .unwrap();
1789
1790 let model = schema.get_model("user_account").unwrap();
1791 assert!(model.get_field("user_id").is_some());
1792 assert!(model.get_field("created_at").is_some());
1793 }
1794
1795 #[test]
1796 fn test_parse_negative_default() {
1797 let schema = parse_schema(
1798 r#"
1799 model Config {
1800 id Int @id
1801 minValue Int @default(-100)
1802 }
1803 "#,
1804 )
1805 .unwrap();
1806
1807 let config = schema.get_model("Config").unwrap();
1808 let min_value = config.get_field("minValue").unwrap();
1809 let attrs = min_value.extract_attributes();
1810 assert!(attrs.default.is_some());
1811 }
1812
1813 #[test]
1814 fn test_parse_float_default() {
1815 let schema = parse_schema(
1816 r#"
1817 model Product {
1818 id Int @id
1819 price Float @default(9.99)
1820 }
1821 "#,
1822 )
1823 .unwrap();
1824
1825 let product = schema.get_model("Product").unwrap();
1826 let price = product.get_field("price").unwrap();
1827 let attrs = price.extract_attributes();
1828 assert!(attrs.default.is_some());
1829 }
1830
1831 #[test]
1834 fn test_parse_simple_server_group() {
1835 let schema = parse_schema(
1836 r#"
1837 serverGroup MainCluster {
1838 server primary {
1839 url = "postgres://localhost/db"
1840 role = "primary"
1841 }
1842 }
1843 "#,
1844 )
1845 .unwrap();
1846
1847 assert_eq!(schema.server_groups.len(), 1);
1848 let cluster = schema.get_server_group("MainCluster").unwrap();
1849 assert_eq!(cluster.servers.len(), 1);
1850 assert!(cluster.servers.contains_key("primary"));
1851 }
1852
1853 #[test]
1854 fn test_parse_server_group_with_multiple_servers() {
1855 let schema = parse_schema(
1856 r#"
1857 serverGroup ReadReplicas {
1858 server primary {
1859 url = "postgres://primary.db.com/app"
1860 role = "primary"
1861 weight = 1
1862 }
1863
1864 server replica1 {
1865 url = "postgres://replica1.db.com/app"
1866 role = "replica"
1867 weight = 2
1868 }
1869
1870 server replica2 {
1871 url = "postgres://replica2.db.com/app"
1872 role = "replica"
1873 weight = 2
1874 }
1875 }
1876 "#,
1877 )
1878 .unwrap();
1879
1880 let cluster = schema.get_server_group("ReadReplicas").unwrap();
1881 assert_eq!(cluster.servers.len(), 3);
1882
1883 let primary = cluster.servers.get("primary").unwrap();
1884 assert_eq!(primary.role(), Some(ServerRole::Primary));
1885 assert_eq!(primary.weight(), Some(1));
1886
1887 let replica1 = cluster.servers.get("replica1").unwrap();
1888 assert_eq!(replica1.role(), Some(ServerRole::Replica));
1889 assert_eq!(replica1.weight(), Some(2));
1890 }
1891
1892 #[test]
1893 fn test_parse_server_group_with_attributes() {
1894 let schema = parse_schema(
1895 r#"
1896 serverGroup ProductionCluster {
1897 @@strategy(ReadReplica)
1898 @@loadBalance(RoundRobin)
1899
1900 server main {
1901 url = "postgres://main/db"
1902 role = "primary"
1903 }
1904 }
1905 "#,
1906 )
1907 .unwrap();
1908
1909 let cluster = schema.get_server_group("ProductionCluster").unwrap();
1910 assert!(cluster.attributes.iter().any(|a| a.name.name == "strategy"));
1911 assert!(
1912 cluster
1913 .attributes
1914 .iter()
1915 .any(|a| a.name.name == "loadBalance")
1916 );
1917 }
1918
1919 #[test]
1920 fn test_parse_server_group_with_env_vars() {
1921 let schema = parse_schema(
1922 r#"
1923 serverGroup EnvCluster {
1924 server db1 {
1925 url = env("PRIMARY_DB_URL")
1926 role = "primary"
1927 }
1928 }
1929 "#,
1930 )
1931 .unwrap();
1932
1933 let cluster = schema.get_server_group("EnvCluster").unwrap();
1934 let server = cluster.servers.get("db1").unwrap();
1935
1936 if let Some(ServerPropertyValue::EnvVar(var)) = server.get_property("url") {
1938 assert_eq!(var, "PRIMARY_DB_URL");
1939 } else {
1940 panic!("Expected env var for url property");
1941 }
1942 }
1943
1944 #[test]
1945 fn test_parse_server_group_with_boolean_property() {
1946 let schema = parse_schema(
1947 r#"
1948 serverGroup TestCluster {
1949 server replica {
1950 url = "postgres://replica/db"
1951 role = "replica"
1952 readOnly = true
1953 }
1954 }
1955 "#,
1956 )
1957 .unwrap();
1958
1959 let cluster = schema.get_server_group("TestCluster").unwrap();
1960 let server = cluster.servers.get("replica").unwrap();
1961 assert!(server.is_read_only());
1962 }
1963
1964 #[test]
1965 fn test_parse_server_group_with_numeric_properties() {
1966 let schema = parse_schema(
1967 r#"
1968 serverGroup NumericCluster {
1969 server db {
1970 url = "postgres://localhost/db"
1971 weight = 5
1972 priority = 1
1973 maxConnections = 100
1974 }
1975 }
1976 "#,
1977 )
1978 .unwrap();
1979
1980 let cluster = schema.get_server_group("NumericCluster").unwrap();
1981 let server = cluster.servers.get("db").unwrap();
1982
1983 assert_eq!(server.weight(), Some(5));
1984 assert_eq!(server.priority(), Some(1));
1985 assert_eq!(server.max_connections(), Some(100));
1986 }
1987
1988 #[test]
1989 fn test_parse_server_group_with_region() {
1990 let schema = parse_schema(
1991 r#"
1992 serverGroup GeoCluster {
1993 server usEast {
1994 url = "postgres://us-east.db.com/app"
1995 role = "replica"
1996 region = "us-east-1"
1997 }
1998
1999 server usWest {
2000 url = "postgres://us-west.db.com/app"
2001 role = "replica"
2002 region = "us-west-2"
2003 }
2004 }
2005 "#,
2006 )
2007 .unwrap();
2008
2009 let cluster = schema.get_server_group("GeoCluster").unwrap();
2010
2011 let us_east = cluster.servers.get("usEast").unwrap();
2012 assert_eq!(us_east.region(), Some("us-east-1"));
2013
2014 let us_west = cluster.servers.get("usWest").unwrap();
2015 assert_eq!(us_west.region(), Some("us-west-2"));
2016
2017 let us_east_servers = cluster.servers_in_region("us-east-1");
2019 assert_eq!(us_east_servers.len(), 1);
2020 }
2021
2022 #[test]
2023 fn test_parse_multiple_server_groups() {
2024 let schema = parse_schema(
2025 r#"
2026 serverGroup Cluster1 {
2027 server db1 {
2028 url = "postgres://db1/app"
2029 }
2030 }
2031
2032 serverGroup Cluster2 {
2033 server db2 {
2034 url = "postgres://db2/app"
2035 }
2036 }
2037
2038 serverGroup Cluster3 {
2039 server db3 {
2040 url = "postgres://db3/app"
2041 }
2042 }
2043 "#,
2044 )
2045 .unwrap();
2046
2047 assert_eq!(schema.server_groups.len(), 3);
2048 assert!(schema.get_server_group("Cluster1").is_some());
2049 assert!(schema.get_server_group("Cluster2").is_some());
2050 assert!(schema.get_server_group("Cluster3").is_some());
2051 }
2052
2053 #[test]
2054 fn test_parse_schema_with_models_and_server_groups() {
2055 let schema = parse_schema(
2056 r#"
2057 model User {
2058 id Int @id @auto
2059 email String @unique
2060 }
2061
2062 serverGroup Database {
2063 @@strategy(ReadReplica)
2064
2065 server primary {
2066 url = env("DATABASE_URL")
2067 role = "primary"
2068 }
2069 }
2070
2071 model Post {
2072 id Int @id @auto
2073 title String
2074 authorId Int
2075 }
2076 "#,
2077 )
2078 .unwrap();
2079
2080 assert_eq!(schema.models.len(), 2);
2081 assert!(schema.get_model("User").is_some());
2082 assert!(schema.get_model("Post").is_some());
2083
2084 assert_eq!(schema.server_groups.len(), 1);
2085 assert!(schema.get_server_group("Database").is_some());
2086 }
2087
2088 #[test]
2089 fn test_parse_server_group_with_health_check() {
2090 let schema = parse_schema(
2091 r#"
2092 serverGroup HealthyCluster {
2093 server monitored {
2094 url = "postgres://localhost/db"
2095 healthCheck = "/health"
2096 }
2097 }
2098 "#,
2099 )
2100 .unwrap();
2101
2102 let cluster = schema.get_server_group("HealthyCluster").unwrap();
2103 let server = cluster.servers.get("monitored").unwrap();
2104 assert_eq!(server.health_check(), Some("/health"));
2105 }
2106
2107 #[test]
2108 fn test_server_group_failover_order() {
2109 let schema = parse_schema(
2110 r#"
2111 serverGroup FailoverCluster {
2112 server db3 {
2113 url = "postgres://db3/app"
2114 priority = 3
2115 }
2116
2117 server db1 {
2118 url = "postgres://db1/app"
2119 priority = 1
2120 }
2121
2122 server db2 {
2123 url = "postgres://db2/app"
2124 priority = 2
2125 }
2126 }
2127 "#,
2128 )
2129 .unwrap();
2130
2131 let cluster = schema.get_server_group("FailoverCluster").unwrap();
2132 let ordered = cluster.failover_order();
2133
2134 assert_eq!(ordered[0].name.name.as_str(), "db1");
2135 assert_eq!(ordered[1].name.name.as_str(), "db2");
2136 assert_eq!(ordered[2].name.name.as_str(), "db3");
2137 }
2138
2139 #[test]
2140 fn test_server_group_names() {
2141 let schema = parse_schema(
2142 r#"
2143 serverGroup Alpha {
2144 server s1 { url = "pg://a" }
2145 }
2146 serverGroup Beta {
2147 server s2 { url = "pg://b" }
2148 }
2149 "#,
2150 )
2151 .unwrap();
2152
2153 let names: Vec<_> = schema.server_group_names().collect();
2154 assert_eq!(names.len(), 2);
2155 assert!(names.contains(&"Alpha"));
2156 assert!(names.contains(&"Beta"));
2157 }
2158
2159 #[test]
2162 fn test_parse_simple_policy() {
2163 let schema = parse_schema(
2164 r#"
2165 policy UserReadOwn on User {
2166 for SELECT
2167 using "id = current_user_id()"
2168 }
2169 "#,
2170 )
2171 .unwrap();
2172
2173 assert_eq!(schema.policies.len(), 1);
2174 let policy = schema.get_policy("UserReadOwn").unwrap();
2175 assert_eq!(policy.name(), "UserReadOwn");
2176 assert_eq!(policy.table(), "User");
2177 assert!(policy.applies_to(PolicyCommand::Select));
2178 assert!(!policy.applies_to(PolicyCommand::Insert));
2179 assert_eq!(policy.using_expr.as_deref(), Some("id = current_user_id()"));
2180 }
2181
2182 #[test]
2183 fn test_parse_policy_with_multiple_commands() {
2184 let schema = parse_schema(
2185 r#"
2186 policy UserModify on User {
2187 for [SELECT, UPDATE, DELETE]
2188 using "id = auth.uid()"
2189 }
2190 "#,
2191 )
2192 .unwrap();
2193
2194 let policy = schema.get_policy("UserModify").unwrap();
2195 assert!(policy.applies_to(PolicyCommand::Select));
2196 assert!(policy.applies_to(PolicyCommand::Update));
2197 assert!(policy.applies_to(PolicyCommand::Delete));
2198 assert!(!policy.applies_to(PolicyCommand::Insert));
2199 }
2200
2201 #[test]
2202 fn test_parse_policy_with_all_command() {
2203 let schema = parse_schema(
2204 r#"
2205 policy UserAll on User {
2206 for ALL
2207 using "true"
2208 }
2209 "#,
2210 )
2211 .unwrap();
2212
2213 let policy = schema.get_policy("UserAll").unwrap();
2214 assert!(policy.applies_to(PolicyCommand::Select));
2215 assert!(policy.applies_to(PolicyCommand::Insert));
2216 assert!(policy.applies_to(PolicyCommand::Update));
2217 assert!(policy.applies_to(PolicyCommand::Delete));
2218 }
2219
2220 #[test]
2221 fn test_parse_policy_with_roles() {
2222 let schema = parse_schema(
2223 r#"
2224 policy AuthenticatedRead on Document {
2225 for SELECT
2226 to authenticated
2227 using "true"
2228 }
2229 "#,
2230 )
2231 .unwrap();
2232
2233 let policy = schema.get_policy("AuthenticatedRead").unwrap();
2234 let roles = policy.effective_roles();
2235 assert!(roles.contains(&"authenticated"));
2236 }
2237
2238 #[test]
2239 fn test_parse_policy_with_multiple_roles() {
2240 let schema = parse_schema(
2241 r#"
2242 policy AdminModerator on Post {
2243 for [UPDATE, DELETE]
2244 to [admin, moderator]
2245 using "true"
2246 }
2247 "#,
2248 )
2249 .unwrap();
2250
2251 let policy = schema.get_policy("AdminModerator").unwrap();
2252 let roles = policy.effective_roles();
2253 assert!(roles.contains(&"admin"));
2254 assert!(roles.contains(&"moderator"));
2255 }
2256
2257 #[test]
2258 fn test_parse_policy_restrictive() {
2259 let schema = parse_schema(
2260 r#"
2261 policy OrgRestriction on Document {
2262 as RESTRICTIVE
2263 for SELECT
2264 using "org_id = current_org_id()"
2265 }
2266 "#,
2267 )
2268 .unwrap();
2269
2270 let policy = schema.get_policy("OrgRestriction").unwrap();
2271 assert!(policy.is_restrictive());
2272 assert!(!policy.is_permissive());
2273 }
2274
2275 #[test]
2276 fn test_parse_policy_permissive_explicit() {
2277 let schema = parse_schema(
2278 r#"
2279 policy Permissive on User {
2280 as PERMISSIVE
2281 for SELECT
2282 using "true"
2283 }
2284 "#,
2285 )
2286 .unwrap();
2287
2288 let policy = schema.get_policy("Permissive").unwrap();
2289 assert!(policy.is_permissive());
2290 }
2291
2292 #[test]
2293 fn test_parse_policy_with_check() {
2294 let schema = parse_schema(
2295 r#"
2296 policy InsertOwn on Post {
2297 for INSERT
2298 to authenticated
2299 check "author_id = current_user_id()"
2300 }
2301 "#,
2302 )
2303 .unwrap();
2304
2305 let policy = schema.get_policy("InsertOwn").unwrap();
2306 assert!(policy.applies_to(PolicyCommand::Insert));
2307 assert_eq!(
2308 policy.check_expr.as_deref(),
2309 Some("author_id = current_user_id()")
2310 );
2311 assert!(policy.using_expr.is_none());
2312 }
2313
2314 #[test]
2315 fn test_parse_policy_with_both_expressions() {
2316 let schema = parse_schema(
2317 r#"
2318 policy UpdateOwn on Post {
2319 for UPDATE
2320 using "author_id = current_user_id()"
2321 check "author_id = current_user_id()"
2322 }
2323 "#,
2324 )
2325 .unwrap();
2326
2327 let policy = schema.get_policy("UpdateOwn").unwrap();
2328 assert!(policy.using_expr.is_some());
2329 assert!(policy.check_expr.is_some());
2330 }
2331
2332 #[test]
2333 fn test_parse_policy_multiline_expression() {
2334 let schema = parse_schema(
2335 r#"
2336 policy ComplexCheck on Document {
2337 for SELECT
2338 using """
2339 (is_public = true)
2340 OR (owner_id = current_user_id())
2341 OR (id IN (SELECT document_id FROM shares WHERE user_id = current_user_id()))
2342 """
2343 }
2344 "#,
2345 )
2346 .unwrap();
2347
2348 let policy = schema.get_policy("ComplexCheck").unwrap();
2349 assert!(policy.using_expr.is_some());
2350 let expr = policy.using_expr.as_ref().unwrap();
2351 assert!(expr.contains("is_public = true"));
2352 assert!(expr.contains("owner_id = current_user_id()"));
2353 assert!(expr.contains("SELECT document_id FROM shares"));
2354 }
2355
2356 #[test]
2357 fn test_parse_multiple_policies() {
2358 let schema = parse_schema(
2359 r#"
2360 policy UserRead on User {
2361 for SELECT
2362 using "true"
2363 }
2364
2365 policy UserInsert on User {
2366 for INSERT
2367 check "id = current_user_id()"
2368 }
2369
2370 policy PostRead on Post {
2371 for SELECT
2372 using "published = true OR author_id = current_user_id()"
2373 }
2374 "#,
2375 )
2376 .unwrap();
2377
2378 assert_eq!(schema.policies.len(), 3);
2379 assert!(schema.get_policy("UserRead").is_some());
2380 assert!(schema.get_policy("UserInsert").is_some());
2381 assert!(schema.get_policy("PostRead").is_some());
2382 }
2383
2384 #[test]
2385 fn test_parse_policy_with_model() {
2386 let schema = parse_schema(
2387 r#"
2388 model User {
2389 id Int @id @auto
2390 email String @unique
2391 }
2392
2393 policy UserReadOwn on User {
2394 for SELECT
2395 to authenticated
2396 using "id = auth.uid()"
2397 }
2398 "#,
2399 )
2400 .unwrap();
2401
2402 assert_eq!(schema.models.len(), 1);
2403 assert_eq!(schema.policies.len(), 1);
2404
2405 let policies = schema.policies_for("User");
2406 assert_eq!(policies.len(), 1);
2407 assert_eq!(policies[0].name(), "UserReadOwn");
2408 }
2409
2410 #[test]
2411 fn test_parse_policies_for_multiple_models() {
2412 let schema = parse_schema(
2413 r#"
2414 policy UserPolicy1 on User {
2415 for SELECT
2416 using "true"
2417 }
2418
2419 policy UserPolicy2 on User {
2420 for INSERT
2421 check "true"
2422 }
2423
2424 policy PostPolicy on Post {
2425 for SELECT
2426 using "true"
2427 }
2428 "#,
2429 )
2430 .unwrap();
2431
2432 assert_eq!(schema.policies_for("User").len(), 2);
2433 assert_eq!(schema.policies_for("Post").len(), 1);
2434 assert!(schema.has_policies("User"));
2435 assert!(schema.has_policies("Post"));
2436 assert!(!schema.has_policies("Comment"));
2437 }
2438
2439 #[test]
2440 fn test_parse_policy_default_all_command() {
2441 let schema = parse_schema(
2442 r#"
2443 policy DefaultAll on User {
2444 using "id = current_user_id()"
2445 }
2446 "#,
2447 )
2448 .unwrap();
2449
2450 let policy = schema.get_policy("DefaultAll").unwrap();
2451 assert!(policy.applies_to(PolicyCommand::All));
2453 }
2454
2455 #[test]
2456 fn test_parse_policy_case_insensitive_keywords() {
2457 let schema = parse_schema(
2458 r#"
2459 policy CaseTest on User {
2460 for select
2461 as permissive
2462 using "true"
2463 }
2464 "#,
2465 )
2466 .unwrap();
2467
2468 let policy = schema.get_policy("CaseTest").unwrap();
2469 assert!(policy.applies_to(PolicyCommand::Select));
2470 assert!(policy.is_permissive());
2471 }
2472
2473 #[test]
2474 fn test_parse_policy_sql_generation() {
2475 let schema = parse_schema(
2476 r#"
2477 model User {
2478 id Int @id
2479
2480 @@map("users")
2481 }
2482
2483 policy ReadOwn on User {
2484 for SELECT
2485 to authenticated
2486 using "id = auth.uid()"
2487 }
2488 "#,
2489 )
2490 .unwrap();
2491
2492 let policy = schema.get_policy("ReadOwn").unwrap();
2493 let sql = policy.to_sql("users");
2494
2495 assert!(sql.contains("CREATE POLICY ReadOwn ON users"));
2496 assert!(sql.contains("FOR SELECT"));
2497 assert!(sql.contains("TO authenticated"));
2498 assert!(sql.contains("USING (id = auth.uid())"));
2499 }
2500
2501 #[test]
2502 fn test_parse_policy_restrictive_sql() {
2503 let schema = parse_schema(
2504 r#"
2505 policy OrgBoundary on Document {
2506 as RESTRICTIVE
2507 for ALL
2508 using "org_id = current_org_id()"
2509 }
2510 "#,
2511 )
2512 .unwrap();
2513
2514 let policy = schema.get_policy("OrgBoundary").unwrap();
2515 let sql = policy.to_sql("documents");
2516
2517 assert!(sql.contains("AS RESTRICTIVE"));
2518 }
2519
2520 #[test]
2521 fn test_parse_policy_with_documentation() {
2522 let schema = parse_schema(
2523 r#"
2524 /// Users can only read their own data
2525 policy UserIsolation on User {
2526 for SELECT
2527 using "id = current_user_id()"
2528 }
2529 "#,
2530 )
2531 .unwrap();
2532
2533 let policy = schema.get_policy("UserIsolation").unwrap();
2534 if let Some(doc) = &policy.documentation {
2535 assert!(doc.text.contains("their own data"));
2536 }
2537 }
2538
2539 #[test]
2540 fn test_parse_complex_rls_schema() {
2541 let schema = parse_schema(
2542 r#"
2543 model Organization {
2544 id Int @id @auto
2545 name String
2546 }
2547
2548 model User {
2549 id Int @id @auto
2550 orgId Int
2551 email String @unique
2552 }
2553
2554 model Document {
2555 id Int @id @auto
2556 title String
2557 ownerId Int
2558 orgId Int
2559 isPublic Boolean @default(false)
2560 }
2561
2562 /// Organization-level isolation
2563 policy OrgIsolation on Document {
2564 as RESTRICTIVE
2565 for ALL
2566 using "org_id = current_setting('app.current_org')::int"
2567 }
2568
2569 /// Users can read public documents
2570 policy PublicRead on Document {
2571 for SELECT
2572 using "is_public = true"
2573 }
2574
2575 /// Users can read their own documents
2576 policy OwnerRead on Document {
2577 for SELECT
2578 to authenticated
2579 using "owner_id = auth.uid()"
2580 }
2581
2582 /// Users can only modify their own documents
2583 policy OwnerModify on Document {
2584 for [UPDATE, DELETE]
2585 to authenticated
2586 using "owner_id = auth.uid()"
2587 check "owner_id = auth.uid()"
2588 }
2589
2590 /// Users can create documents in their org
2591 policy OrgInsert on Document {
2592 for INSERT
2593 to authenticated
2594 check "org_id = current_setting('app.current_org')::int"
2595 }
2596 "#,
2597 )
2598 .unwrap();
2599
2600 assert_eq!(schema.models.len(), 3);
2601 assert_eq!(schema.policies.len(), 5);
2602
2603 let org_iso = schema.get_policy("OrgIsolation").unwrap();
2605 assert!(org_iso.is_restrictive());
2606
2607 let doc_policies = schema.policies_for("Document");
2609 assert_eq!(doc_policies.len(), 5);
2610 }
2611
2612 #[test]
2615 fn test_parse_policy_with_mssql_schema() {
2616 let schema = parse_schema(
2617 r#"
2618 policy UserFilter on User {
2619 for SELECT
2620 using "UserId = @UserId"
2621 mssqlSchema "RLS"
2622 }
2623 "#,
2624 )
2625 .unwrap();
2626
2627 let policy = schema.get_policy("UserFilter").unwrap();
2628 assert_eq!(policy.mssql_schema(), "RLS");
2629 }
2630
2631 #[test]
2632 fn test_parse_policy_with_mssql_block_single() {
2633 let schema = parse_schema(
2634 r#"
2635 policy UserInsert on User {
2636 for INSERT
2637 check "UserId = @UserId"
2638 mssqlBlock AFTER_INSERT
2639 }
2640 "#,
2641 )
2642 .unwrap();
2643
2644 let policy = schema.get_policy("UserInsert").unwrap();
2645 assert_eq!(policy.mssql_block_operations.len(), 1);
2646 assert_eq!(
2647 policy.mssql_block_operations[0],
2648 MssqlBlockOperation::AfterInsert
2649 );
2650 }
2651
2652 #[test]
2653 fn test_parse_policy_with_mssql_block_list() {
2654 let schema = parse_schema(
2655 r#"
2656 policy UserModify on User {
2657 for [INSERT, UPDATE, DELETE]
2658 check "UserId = @UserId"
2659 mssqlBlock [AFTER_INSERT, AFTER_UPDATE, BEFORE_DELETE]
2660 }
2661 "#,
2662 )
2663 .unwrap();
2664
2665 let policy = schema.get_policy("UserModify").unwrap();
2666 assert_eq!(policy.mssql_block_operations.len(), 3);
2667 assert!(
2668 policy
2669 .mssql_block_operations
2670 .contains(&MssqlBlockOperation::AfterInsert)
2671 );
2672 assert!(
2673 policy
2674 .mssql_block_operations
2675 .contains(&MssqlBlockOperation::AfterUpdate)
2676 );
2677 assert!(
2678 policy
2679 .mssql_block_operations
2680 .contains(&MssqlBlockOperation::BeforeDelete)
2681 );
2682 }
2683
2684 #[test]
2685 fn test_parse_policy_full_mssql_config() {
2686 let schema = parse_schema(
2687 r#"
2688 policy TenantIsolation on Order {
2689 for ALL
2690 using "TenantId = @TenantId"
2691 check "TenantId = @TenantId"
2692 mssqlSchema "MultiTenant"
2693 mssqlBlock [AFTER_INSERT, BEFORE_UPDATE, AFTER_UPDATE, BEFORE_DELETE]
2694 }
2695 "#,
2696 )
2697 .unwrap();
2698
2699 let policy = schema.get_policy("TenantIsolation").unwrap();
2700
2701 assert!(policy.applies_to(PolicyCommand::All));
2703 assert!(policy.using_expr.is_some());
2704 assert!(policy.check_expr.is_some());
2705
2706 assert_eq!(policy.mssql_schema(), "MultiTenant");
2708 assert_eq!(policy.mssql_block_operations.len(), 4);
2709
2710 let mssql = policy.to_mssql_sql("dbo.Orders", "TenantId");
2712 assert!(mssql.schema_sql.contains("MultiTenant"));
2713 assert!(mssql.function_sql.contains("fn_TenantIsolation_predicate"));
2714 }
2715
2716 #[test]
2717 fn test_parse_policy_mssql_block_case_variants() {
2718 let schema = parse_schema(
2720 r#"
2721 policy Test1 on User {
2722 for INSERT
2723 check "true"
2724 mssqlBlock after_insert
2725 }
2726 "#,
2727 )
2728 .unwrap();
2729
2730 let policy = schema.get_policy("Test1").unwrap();
2731 assert_eq!(policy.mssql_block_operations.len(), 1);
2732 assert_eq!(
2733 policy.mssql_block_operations[0],
2734 MssqlBlockOperation::AfterInsert
2735 );
2736 }
2737
2738 #[test]
2739 fn test_parse_mixed_postgres_mssql_schema() {
2740 let schema = parse_schema(
2741 r#"
2742 model User {
2743 id Int @id @auto
2744 email String @unique
2745 }
2746
2747 // PostgreSQL-style policy (works on both, MSSQL uses defaults)
2748 policy UserReadOwn on User {
2749 for SELECT
2750 to authenticated
2751 using "id = current_user_id()"
2752 }
2753
2754 // MSSQL-optimized policy with explicit settings
2755 policy UserModifyOwn on User {
2756 for [INSERT, UPDATE, DELETE]
2757 to authenticated
2758 using "id = current_user_id()"
2759 check "id = current_user_id()"
2760 mssqlSchema "Security"
2761 mssqlBlock [AFTER_INSERT, BEFORE_UPDATE, AFTER_UPDATE, BEFORE_DELETE]
2762 }
2763 "#,
2764 )
2765 .unwrap();
2766
2767 assert_eq!(schema.policies.len(), 2);
2768
2769 let read_policy = schema.get_policy("UserReadOwn").unwrap();
2771 assert_eq!(read_policy.mssql_schema(), "Security"); assert!(read_policy.mssql_block_operations.is_empty()); let modify_policy = schema.get_policy("UserModifyOwn").unwrap();
2776 assert_eq!(modify_policy.mssql_schema(), "Security");
2777 assert_eq!(modify_policy.mssql_block_operations.len(), 4);
2778
2779 let pg_sql = read_policy.to_postgres_sql("users");
2781 assert!(pg_sql.contains("CREATE POLICY UserReadOwn ON users"));
2782
2783 let mssql = modify_policy.to_mssql_sql("dbo.Users", "id");
2785 assert!(mssql.policy_sql.contains("Security.UserModifyOwn"));
2786 }
2787}