Skip to main content

spikard_cli/codegen/protobuf/
spec_parser.rs

1//! Protobuf (.proto) specification parsing and extraction.
2//!
3//! This module handles parsing Protocol Buffer specifications (proto3 syntax only)
4//! and extracting structured data for code generation, including messages, services,
5//! enums, and field definitions.
6
7use anyhow::{Context, Result, anyhow, bail};
8use serde::{Deserialize, Serialize};
9use std::collections::{HashMap, HashSet};
10use std::fs;
11use std::path::{Path, PathBuf};
12
13/// Parsed Protobuf schema representation
14#[derive(Debug, Clone, Serialize, Deserialize)]
15pub struct ProtobufSchema {
16    /// Package name (e.g., "com.example.service")
17    pub package: Option<String>,
18    /// Map of message names to their definitions
19    pub messages: HashMap<String, MessageDef>,
20    /// Map of service names to their definitions
21    pub services: HashMap<String, ServiceDef>,
22    /// Map of enum names to their definitions
23    pub enums: HashMap<String, EnumDef>,
24    /// List of imported proto files
25    pub imports: Vec<String>,
26    /// Proto file syntax version (enforced to be "proto3")
27    pub syntax: String,
28    /// Schema description/comments
29    pub description: Option<String>,
30}
31
32/// Protobuf message definition
33#[derive(Debug, Clone, Serialize, Deserialize)]
34pub struct MessageDef {
35    /// Message name
36    pub name: String,
37    /// Message fields
38    pub fields: Vec<FieldDef>,
39    /// Nested message definitions
40    pub nested_messages: HashMap<String, Self>,
41    /// Nested enum definitions
42    pub nested_enums: HashMap<String, EnumDef>,
43    /// Message description from comments
44    pub description: Option<String>,
45}
46
47/// Protobuf service definition
48#[derive(Debug, Clone, Serialize, Deserialize)]
49pub struct ServiceDef {
50    /// Service name
51    pub name: String,
52    /// Service methods/RPCs
53    pub methods: Vec<MethodDef>,
54    /// Service description from comments
55    pub description: Option<String>,
56}
57
58/// Protobuf RPC method definition
59#[derive(Debug, Clone, Serialize, Deserialize)]
60pub struct MethodDef {
61    /// Method name
62    pub name: String,
63    /// Input message type name
64    pub input_type: String,
65    /// Output message type name
66    pub output_type: String,
67    /// Whether input is a stream
68    pub input_streaming: bool,
69    /// Whether output is a stream
70    pub output_streaming: bool,
71    /// Method description from comments
72    pub description: Option<String>,
73}
74
75/// Protobuf field definition
76#[derive(Debug, Clone, Serialize, Deserialize)]
77pub struct FieldDef {
78    /// Field name
79    pub name: String,
80    /// Field number (1-536870911)
81    pub number: u32,
82    /// Field type
83    pub field_type: ProtoType,
84    /// Field label (optional, repeated, or neither for required)
85    pub label: FieldLabel,
86    /// Default value (if applicable)
87    pub default_value: Option<String>,
88    /// Field description from comments
89    pub description: Option<String>,
90}
91
92/// Protocol Buffer field label
93#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
94pub enum FieldLabel {
95    /// No label (proto3 default: optional for scalars, required for messages)
96    None,
97    /// Repeated field (becomes a list)
98    Repeated,
99    /// Optional field (may be unset)
100    Optional,
101}
102
103/// Protobuf enum definition
104#[derive(Debug, Clone, Serialize, Deserialize)]
105pub struct EnumDef {
106    /// Enum name
107    pub name: String,
108    /// Enum values
109    pub values: Vec<EnumValue>,
110    /// Enum description from comments
111    pub description: Option<String>,
112}
113
114/// Protobuf enum value
115#[derive(Debug, Clone, Serialize, Deserialize)]
116pub struct EnumValue {
117    /// Value name
118    pub name: String,
119    /// Numeric value
120    pub number: i32,
121    /// Value description from comments
122    pub description: Option<String>,
123}
124
125/// Protocol Buffer type enumeration
126#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
127pub enum ProtoType {
128    Double,
129    Float,
130    Int32,
131    Int64,
132    Uint32,
133    Uint64,
134    Sint32,
135    Sint64,
136    Fixed32,
137    Fixed64,
138    Sfixed32,
139    Sfixed64,
140    Bool,
141    String,
142    Bytes,
143    Message(String),
144    Enum(String),
145}
146
147impl ProtoType {
148    /// Get the string representation of a proto type
149    #[must_use]
150    pub fn as_str(&self) -> String {
151        match self {
152            Self::Double => "double".to_string(),
153            Self::Float => "float".to_string(),
154            Self::Int32 => "int32".to_string(),
155            Self::Int64 => "int64".to_string(),
156            Self::Uint32 => "uint32".to_string(),
157            Self::Uint64 => "uint64".to_string(),
158            Self::Sint32 => "sint32".to_string(),
159            Self::Sint64 => "sint64".to_string(),
160            Self::Fixed32 => "fixed32".to_string(),
161            Self::Fixed64 => "fixed64".to_string(),
162            Self::Sfixed32 => "sfixed32".to_string(),
163            Self::Sfixed64 => "sfixed64".to_string(),
164            Self::Bool => "bool".to_string(),
165            Self::String => "string".to_string(),
166            Self::Bytes => "bytes".to_string(),
167            Self::Message(name) => name.clone(),
168            Self::Enum(name) => name.clone(),
169        }
170    }
171}
172
173/// Parse a Protobuf schema from a .proto file
174///
175/// # Arguments
176/// * `path` - Path to .proto file
177///
178/// # Returns
179/// Parsed `ProtobufSchema` or error (rejects proto2 syntax)
180pub fn parse_proto_schema(path: &Path) -> Result<ProtobufSchema> {
181    let content = fs::read_to_string(path).with_context(|| format!("Failed to read proto file: {}", path.display()))?;
182
183    parse_proto_schema_string(&content).with_context(|| format!("Failed to parse proto schema from {}", path.display()))
184}
185
186/// Parse a Protobuf schema from a .proto file and recursively merge import dependencies.
187///
188/// Imported files are resolved relative to the source file first and then against
189/// any additional include paths supplied by the caller.
190pub fn parse_proto_schema_with_includes(path: &Path, include_paths: &[PathBuf]) -> Result<ProtobufSchema> {
191    let mut visited = HashSet::new();
192    parse_proto_schema_recursive(path, include_paths, &mut visited)
193}
194
195/// Parse a Protobuf schema from a string
196pub fn parse_proto_schema_string(content: &str) -> Result<ProtobufSchema> {
197    let mut schema = ProtobufSchema {
198        package: None,
199        messages: HashMap::new(),
200        services: HashMap::new(),
201        enums: HashMap::new(),
202        imports: Vec::new(),
203        syntax: String::new(),
204        description: None,
205    };
206
207    schema.syntax = extract_syntax_declaration(content).unwrap_or_else(|| "proto3".to_string());
208
209    if schema.syntax != "proto3" {
210        return Err(anyhow!(
211            "Only proto3 syntax is supported. Found: {}\n\
212             Please convert your proto file to proto3 syntax or use proto3-compatible definitions.\n\
213             See: https://developers.google.com/protocol-buffers/docs/proto3",
214            schema.syntax
215        ));
216    }
217
218    schema.package = extract_package_name(content);
219
220    schema.imports = extract_imports(content);
221
222    parse_top_level_definitions(content, &mut schema)?;
223
224    Ok(schema)
225}
226
227fn parse_proto_schema_recursive(
228    path: &Path,
229    include_paths: &[PathBuf],
230    visited: &mut HashSet<PathBuf>,
231) -> Result<ProtobufSchema> {
232    let visit_key = canonical_or_original(path);
233    if !visited.insert(visit_key) {
234        return Ok(ProtobufSchema {
235            package: None,
236            messages: HashMap::new(),
237            services: HashMap::new(),
238            enums: HashMap::new(),
239            imports: Vec::new(),
240            syntax: "proto3".to_string(),
241            description: None,
242        });
243    }
244
245    let content = fs::read_to_string(path).with_context(|| format!("Failed to read proto file: {}", path.display()))?;
246    let mut schema = parse_proto_schema_string(&content)
247        .with_context(|| format!("Failed to parse proto schema from {}", path.display()))?;
248
249    for import in schema.imports.clone() {
250        let Some(import_path) = resolve_import_path(path, &import, include_paths) else {
251            continue;
252        };
253        let imported_schema = parse_proto_schema_recursive(&import_path, include_paths, visited)?;
254        merge_schema(&mut schema, imported_schema)?;
255    }
256
257    Ok(schema)
258}
259
260fn resolve_import_path(path: &Path, import: &str, include_paths: &[PathBuf]) -> Option<PathBuf> {
261    let mut relative_candidates = path
262        .parent()
263        .into_iter()
264        .map(|parent| parent.join(import))
265        .chain(include_paths.iter().map(|include| include.join(import)));
266
267    relative_candidates
268        .find(|candidate| candidate.is_file())
269        .map(|candidate| canonical_or_original(&candidate))
270}
271
272fn canonical_or_original(path: &Path) -> PathBuf {
273    fs::canonicalize(path).unwrap_or_else(|_| path.to_path_buf())
274}
275
276fn merge_schema(target: &mut ProtobufSchema, imported: ProtobufSchema) -> Result<()> {
277    merge_named_defs("message", &mut target.messages, imported.messages)?;
278    merge_named_defs("enum", &mut target.enums, imported.enums)?;
279    merge_named_defs("service", &mut target.services, imported.services)?;
280
281    for import in imported.imports {
282        if !target.imports.contains(&import) {
283            target.imports.push(import);
284        }
285    }
286
287    Ok(())
288}
289
290fn merge_named_defs<T>(kind: &str, target: &mut HashMap<String, T>, source: HashMap<String, T>) -> Result<()> {
291    for (name, def) in source {
292        if target.contains_key(&name) {
293            bail!("Duplicate {kind} definition found while resolving imports: {name}");
294        }
295        target.insert(name, def);
296    }
297    Ok(())
298}
299
300/// Helper function to extract syntax declaration from proto content
301fn extract_syntax_declaration(content: &str) -> Option<String> {
302    for line in content.lines() {
303        let trimmed = line.trim();
304        if trimmed.starts_with("syntax") {
305            let quote_start = trimmed.find('"')?;
306            let remaining = &trimmed[quote_start + 1..];
307            let quote_end = remaining.find('"')?;
308            return Some(remaining[..quote_end].to_string());
309        }
310    }
311    None
312}
313
314/// Helper function to extract package name from proto content
315fn extract_package_name(content: &str) -> Option<String> {
316    for line in content.lines() {
317        let trimmed = line.trim();
318        if trimmed.starts_with("package") && !trimmed.starts_with("package ") {
319            continue;
320        }
321        if let Some(package_part) = trimmed.strip_prefix("package ") {
322            let semicolon_pos = package_part.find(';')?;
323            let package_name = package_part[..semicolon_pos].trim();
324            return Some(package_name.to_string());
325        }
326    }
327    None
328}
329
330/// Helper function to extract imports from proto content
331fn extract_imports(content: &str) -> Vec<String> {
332    let mut imports = Vec::new();
333    for line in content.lines() {
334        let trimmed = line.trim();
335        if trimmed.starts_with("import ") && trimmed.contains('"') {
336            if let Some(quote_start) = trimmed.find('"') {
337                let remaining = &trimmed[quote_start + 1..];
338                if let Some(quote_end) = remaining.find('"') {
339                    imports.push(remaining[..quote_end].to_string());
340                }
341            }
342        }
343    }
344    imports
345}
346
347fn parse_top_level_definitions(content: &str, schema: &mut ProtobufSchema) -> Result<()> {
348    let lines: Vec<&str> = content.lines().collect();
349    let mut index = 0;
350    let mut pending_comment: Vec<String> = Vec::new();
351
352    while index < lines.len() {
353        let trimmed = strip_inline_comment(lines[index]).trim();
354
355        if trimmed.is_empty() {
356            if !pending_comment.is_empty() {
357                pending_comment.clear();
358            }
359            index += 1;
360            continue;
361        }
362
363        if let Some(comment) = lines[index].trim().strip_prefix("//") {
364            pending_comment.push(comment.trim().to_string());
365            index += 1;
366            continue;
367        }
368
369        if trimmed.starts_with("message ") {
370            let (message, next_index) = parse_message_block(&lines, index, take_comment(&mut pending_comment))?;
371            schema.messages.insert(message.name.clone(), message);
372            index = next_index;
373            continue;
374        }
375
376        if trimmed.starts_with("enum ") {
377            let (enum_def, next_index) = parse_enum_block(&lines, index, take_comment(&mut pending_comment))?;
378            schema.enums.insert(enum_def.name.clone(), enum_def);
379            index = next_index;
380            continue;
381        }
382
383        if trimmed.starts_with("service ") {
384            let (service, next_index) = parse_service_block(&lines, index, take_comment(&mut pending_comment))?;
385            schema.services.insert(service.name.clone(), service);
386            index = next_index;
387            continue;
388        }
389
390        pending_comment.clear();
391        index += 1;
392    }
393
394    Ok(())
395}
396
397/// Split the inline body of a block header into `;`-terminated statements and report the
398/// net brace depth the header leaves open. A single-line block (`message X { f = 1; }`)
399/// closes on its own line, so it must contribute its statements here and return depth 0 —
400/// otherwise the caller counts only the opening brace, stays "open", and swallows every
401/// following definition until braces happen to net out. Statements are re-terminated with
402/// `;` because the field/value parsers require it. ~keep
403fn header_inline_body(header: &str) -> (Vec<String>, usize) {
404    let depth = header.matches('{').count().saturating_sub(header.matches('}').count());
405    let mut items = Vec::new();
406    if let Some(pos) = header.find('{') {
407        let mut inline = header[pos + 1..].trim();
408        if let Some(stripped) = inline.strip_suffix('}') {
409            inline = stripped.trim();
410        }
411        for part in inline.split(';') {
412            let part = part.trim();
413            if !part.is_empty() {
414                items.push(format!("{part};"));
415            }
416        }
417    }
418    (items, depth)
419}
420
421fn parse_message_block(lines: &[&str], start: usize, description: Option<String>) -> Result<(MessageDef, usize)> {
422    let header = strip_inline_comment(lines[start]).trim();
423    let name = extract_block_name(header, "message")
424        .ok_or_else(|| anyhow!("Invalid message declaration: {}", lines[start].trim()))?;
425
426    let mut message = MessageDef {
427        name,
428        fields: Vec::new(),
429        nested_messages: HashMap::new(),
430        nested_enums: HashMap::new(),
431        description,
432    };
433
434    let index = start + 1;
435    let (inline_fields, depth) = header_inline_body(header);
436    for field_line in &inline_fields {
437        if !field_line.starts_with("message ") && !field_line.starts_with("enum ") {
438            if let Some(field) = parse_field(field_line, None)? {
439                message.fields.push(field);
440            }
441        }
442    }
443    if depth == 0 {
444        return Ok((message, index));
445    }
446
447    let mut index = index;
448    let mut depth = depth;
449    let mut pending_comment: Vec<String> = Vec::new();
450
451    while index < lines.len() {
452        let raw_line = lines[index];
453        let line = strip_inline_comment(raw_line);
454        let trimmed = line.trim();
455
456        if trimmed.starts_with("//") {
457            if let Some(comment) = raw_line.trim().strip_prefix("//") {
458                pending_comment.push(comment.trim().to_string());
459            }
460            index += 1;
461            continue;
462        }
463
464        let opens = trimmed.matches('{').count();
465        let closes = trimmed.matches('}').count();
466
467        if depth == 1 && !trimmed.is_empty() && !trimmed.starts_with("message ") && !trimmed.starts_with("enum ") {
468            if let Some(field) = parse_field(trimmed, take_comment(&mut pending_comment))? {
469                message.fields.push(field);
470            }
471        }
472
473        depth += opens;
474        depth = depth.saturating_sub(closes);
475        index += 1;
476
477        if depth == 0 {
478            break;
479        }
480    }
481
482    Ok((message, index))
483}
484
485fn parse_enum_block(lines: &[&str], start: usize, description: Option<String>) -> Result<(EnumDef, usize)> {
486    let header = strip_inline_comment(lines[start]).trim();
487    let name = extract_block_name(header, "enum")
488        .ok_or_else(|| anyhow!("Invalid enum declaration: {}", lines[start].trim()))?;
489
490    let mut enum_def = EnumDef {
491        name,
492        values: Vec::new(),
493        description,
494    };
495
496    let index = start + 1;
497    let (inline_values, depth) = header_inline_body(header);
498    for value_line in &inline_values {
499        if value_line.contains('=') {
500            if let Some(value) = parse_enum_value(value_line, None)? {
501                enum_def.values.push(value);
502            }
503        }
504    }
505    if depth == 0 {
506        return Ok((enum_def, index));
507    }
508
509    let mut index = index;
510    let mut depth = depth;
511    let mut pending_comment: Vec<String> = Vec::new();
512
513    while index < lines.len() {
514        let raw_line = lines[index];
515        let line = strip_inline_comment(raw_line);
516        let trimmed = line.trim();
517
518        if trimmed.starts_with("//") {
519            if let Some(comment) = raw_line.trim().strip_prefix("//") {
520                pending_comment.push(comment.trim().to_string());
521            }
522            index += 1;
523            continue;
524        }
525
526        let opens = trimmed.matches('{').count();
527        let closes = trimmed.matches('}').count();
528
529        if depth == 1 && trimmed.contains('=') && trimmed.ends_with(';') {
530            if let Some(value) = parse_enum_value(trimmed, take_comment(&mut pending_comment))? {
531                enum_def.values.push(value);
532            }
533        }
534
535        depth += opens;
536        depth = depth.saturating_sub(closes);
537        index += 1;
538
539        if depth == 0 {
540            break;
541        }
542    }
543
544    Ok((enum_def, index))
545}
546
547fn parse_service_block(lines: &[&str], start: usize, description: Option<String>) -> Result<(ServiceDef, usize)> {
548    let header = strip_inline_comment(lines[start]).trim();
549    let name = extract_block_name(header, "service")
550        .ok_or_else(|| anyhow!("Invalid service declaration: {}", lines[start].trim()))?;
551
552    let mut service = ServiceDef {
553        name,
554        methods: Vec::new(),
555        description,
556    };
557
558    let mut index = start + 1;
559    let mut depth = header.matches('{').count().saturating_sub(header.matches('}').count());
560    if depth == 0 {
561        return Ok((service, index));
562    }
563    let mut pending_comment: Vec<String> = Vec::new();
564
565    while index < lines.len() {
566        let raw_line = lines[index];
567        let line = strip_inline_comment(raw_line);
568        let trimmed = line.trim();
569
570        if trimmed.starts_with("//") {
571            if let Some(comment) = raw_line.trim().strip_prefix("//") {
572                pending_comment.push(comment.trim().to_string());
573            }
574            index += 1;
575            continue;
576        }
577
578        let opens = trimmed.matches('{').count();
579        let closes = trimmed.matches('}').count();
580
581        if depth == 1 && trimmed.starts_with("rpc ") {
582            if let Some(method) = parse_rpc_method(trimmed, take_comment(&mut pending_comment))? {
583                service.methods.push(method);
584            }
585        }
586
587        depth += opens;
588        depth = depth.saturating_sub(closes);
589        index += 1;
590
591        if depth == 0 {
592            break;
593        }
594    }
595
596    Ok((service, index))
597}
598
599fn parse_field(line: &str, description: Option<String>) -> Result<Option<FieldDef>> {
600    if !line.ends_with(';') || line.starts_with("option ") || line.starts_with("reserved ") {
601        return Ok(None);
602    }
603
604    let without_semicolon = line.trim_end_matches(';');
605    let declaration = without_semicolon.split('[').next().unwrap_or(without_semicolon).trim();
606    let parts: Vec<&str> = declaration.split_whitespace().collect();
607
608    if parts.len() < 4 {
609        return Ok(None);
610    }
611
612    let (label, type_index) = match parts[0] {
613        "repeated" => (FieldLabel::Repeated, 1),
614        "optional" => (FieldLabel::Optional, 1),
615        _ => (FieldLabel::None, 0),
616    };
617
618    if parts.len() <= type_index + 2 {
619        return Ok(None);
620    }
621
622    let field_type = parse_proto_type(parts[type_index]);
623    let field_name = parts[type_index + 1].to_string();
624    let number = parts[type_index + 3]
625        .parse::<u32>()
626        .with_context(|| format!("Invalid field number in line: {line}"))?;
627
628    Ok(Some(FieldDef {
629        name: field_name,
630        number,
631        field_type,
632        label,
633        default_value: None,
634        description,
635    }))
636}
637
638fn parse_enum_value(line: &str, description: Option<String>) -> Result<Option<EnumValue>> {
639    let without_semicolon = line.trim_end_matches(';').trim();
640    let (name, number) = without_semicolon
641        .split_once('=')
642        .ok_or_else(|| anyhow!("Invalid enum value declaration: {line}"))?;
643
644    Ok(Some(EnumValue {
645        name: name.trim().to_string(),
646        number: number
647            .trim()
648            .parse::<i32>()
649            .with_context(|| format!("Invalid enum value number in line: {line}"))?,
650        description,
651    }))
652}
653
654fn parse_rpc_method(line: &str, description: Option<String>) -> Result<Option<MethodDef>> {
655    let without_semicolon = line.trim_end_matches(';').trim();
656    let after_rpc = without_semicolon
657        .strip_prefix("rpc ")
658        .ok_or_else(|| anyhow!("Invalid RPC declaration: {line}"))?;
659    let method_name_end = after_rpc
660        .find('(')
661        .ok_or_else(|| anyhow!("Invalid RPC declaration: {line}"))?;
662    let method_name = after_rpc[..method_name_end].trim().to_string();
663    let rest = &after_rpc[method_name_end + 1..];
664    let request_end = rest
665        .find(')')
666        .ok_or_else(|| anyhow!("Invalid RPC request declaration: {line}"))?;
667    let request_decl = rest[..request_end].trim();
668    let after_request = rest[request_end + 1..].trim();
669    let returns_decl = after_request
670        .strip_prefix("returns")
671        .ok_or_else(|| anyhow!("Invalid RPC returns declaration: {line}"))?
672        .trim();
673    let returns_decl = returns_decl
674        .strip_prefix('(')
675        .ok_or_else(|| anyhow!("Invalid RPC returns declaration: {line}"))?;
676    let response_end = returns_decl
677        .find(')')
678        .ok_or_else(|| anyhow!("Invalid RPC returns declaration: {line}"))?;
679    let response_decl = returns_decl[..response_end].trim();
680
681    let (input_streaming, input_type) = parse_streaming_type(request_decl);
682    let (output_streaming, output_type) = parse_streaming_type(response_decl);
683
684    Ok(Some(MethodDef {
685        name: method_name,
686        input_type,
687        output_type,
688        input_streaming,
689        output_streaming,
690        description,
691    }))
692}
693
694fn parse_streaming_type(declaration: &str) -> (bool, String) {
695    if let Some(rest) = declaration.strip_prefix("stream ") {
696        (true, rest.trim().to_string())
697    } else {
698        (false, declaration.trim().to_string())
699    }
700}
701
702fn extract_block_name(header: &str, keyword: &str) -> Option<String> {
703    header
704        .strip_prefix(keyword)?
705        .trim()
706        .strip_suffix('{')
707        .unwrap_or_else(|| header.strip_prefix(keyword).unwrap().trim())
708        .split_whitespace()
709        .next()
710        .map(std::string::ToString::to_string)
711}
712
713fn strip_inline_comment(line: &str) -> &str {
714    if let Some((before, _)) = line.split_once("//") {
715        before
716    } else {
717        line
718    }
719}
720
721fn take_comment(pending_comment: &mut Vec<String>) -> Option<String> {
722    if pending_comment.is_empty() {
723        None
724    } else {
725        let comment = pending_comment.join(" ");
726        pending_comment.clear();
727        Some(comment)
728    }
729}
730
731/// Helper function to parse type name from proto syntax
732#[allow(dead_code)]
733fn parse_proto_type(type_str: &str) -> ProtoType {
734    match type_str {
735        "double" => ProtoType::Double,
736        "float" => ProtoType::Float,
737        "int32" => ProtoType::Int32,
738        "int64" => ProtoType::Int64,
739        "uint32" => ProtoType::Uint32,
740        "uint64" => ProtoType::Uint64,
741        "sint32" => ProtoType::Sint32,
742        "sint64" => ProtoType::Sint64,
743        "fixed32" => ProtoType::Fixed32,
744        "fixed64" => ProtoType::Fixed64,
745        "sfixed32" => ProtoType::Sfixed32,
746        "sfixed64" => ProtoType::Sfixed64,
747        "bool" => ProtoType::Bool,
748        "string" => ProtoType::String,
749        "bytes" => ProtoType::Bytes,
750        _ => ProtoType::Message(type_str.to_string()),
751    }
752}
753
754#[cfg(test)]
755mod tests {
756    use super::*;
757    use tempfile::tempdir;
758
759    #[test]
760    fn test_parse_simple_proto3_schema() {
761        let proto = r#"syntax = "proto3";
762
763package example;
764
765message User {
766  string id = 1;
767  string name = 2;
768  string email = 3;
769}
770"#;
771
772        let schema = parse_proto_schema_string(proto).expect("Failed to parse proto");
773        assert_eq!(schema.syntax, "proto3");
774        assert_eq!(schema.package, Some("example".to_string()));
775        let user = schema.messages.get("User").expect("message should be parsed");
776        assert_eq!(user.fields.len(), 3);
777        assert_eq!(user.fields[0].name, "id");
778    }
779
780    #[test]
781    fn test_parse_proto_with_imports() {
782        let proto = r#"syntax = "proto3";
783
784import "google/protobuf/timestamp.proto";
785import "other.proto";
786
787package example;
788"#;
789
790        let schema = parse_proto_schema_string(proto).expect("Failed to parse proto");
791        assert_eq!(schema.imports.len(), 2);
792        assert!(schema.imports.contains(&"google/protobuf/timestamp.proto".to_string()));
793        assert!(schema.imports.contains(&"other.proto".to_string()));
794    }
795
796    #[test]
797    fn test_parse_proto_schema_with_includes_merges_imported_messages() {
798        let temp_dir = tempdir().expect("temp dir");
799        let shared_dir = temp_dir.path().join("common");
800        fs::create_dir_all(&shared_dir).expect("create include dir");
801
802        let shared_proto = shared_dir.join("types.proto");
803        fs::write(
804            &shared_proto,
805            r#"syntax = "proto3";
806
807package common;
808
809message SharedType {
810  string id = 1;
811}
812"#,
813        )
814        .expect("write shared proto");
815
816        let root_proto = temp_dir.path().join("service.proto");
817        fs::write(
818            &root_proto,
819            r#"syntax = "proto3";
820
821import "common/types.proto";
822
823package example;
824
825message UsesShared {
826  SharedType shared = 1;
827}
828"#,
829        )
830        .expect("write root proto");
831
832        let schema = parse_proto_schema_with_includes(&root_proto, &[temp_dir.path().to_path_buf()])
833            .expect("schema should resolve imports");
834
835        assert!(schema.messages.contains_key("UsesShared"));
836        assert!(schema.messages.contains_key("SharedType"));
837        assert!(schema.imports.contains(&"common/types.proto".to_string()));
838    }
839
840    #[test]
841    fn test_reject_proto2_syntax() {
842        let proto = r#"syntax = "proto2";
843
844package example;
845
846message User {
847  required string id = 1;
848}
849"#;
850
851        let result = parse_proto_schema_string(proto);
852        assert!(result.is_err());
853        let error_msg = format!("{}", result.unwrap_err());
854        assert!(error_msg.contains("Only proto3 syntax is supported"));
855        assert!(error_msg.contains("proto2"));
856    }
857
858    #[test]
859    fn test_parse_proto_type_scalars() {
860        assert_eq!(parse_proto_type("double"), ProtoType::Double);
861        assert_eq!(parse_proto_type("float"), ProtoType::Float);
862        assert_eq!(parse_proto_type("int32"), ProtoType::Int32);
863        assert_eq!(parse_proto_type("int64"), ProtoType::Int64);
864        assert_eq!(parse_proto_type("bool"), ProtoType::Bool);
865        assert_eq!(parse_proto_type("string"), ProtoType::String);
866        assert_eq!(parse_proto_type("bytes"), ProtoType::Bytes);
867    }
868
869    #[test]
870    fn test_parse_proto_type_message() {
871        match parse_proto_type("User") {
872            ProtoType::Message(name) => assert_eq!(name, "User"),
873            _ => panic!("Expected Message type"),
874        }
875    }
876
877    #[test]
878    fn test_parse_service_and_enum() {
879        let proto = r#"syntax = "proto3";
880
881package example;
882
883enum Status {
884  STATUS_UNKNOWN = 0;
885  STATUS_ACTIVE = 1;
886}
887
888service UserService {
889  rpc GetUser (GetUserRequest) returns (User);
890  rpc ListUsers (ListUsersRequest) returns (stream User);
891}
892"#;
893
894        let schema = parse_proto_schema_string(proto).expect("Failed to parse proto");
895        let status = schema.enums.get("Status").expect("enum should be parsed");
896        assert_eq!(status.values.len(), 2);
897
898        let service = schema.services.get("UserService").expect("service should be parsed");
899        assert_eq!(service.methods.len(), 2);
900        assert_eq!(service.methods[0].name, "GetUser");
901        assert!(service.methods[1].output_streaming);
902    }
903
904    #[test]
905    fn test_proto_type_as_str() {
906        assert_eq!(ProtoType::Double.as_str(), "double");
907        assert_eq!(ProtoType::String.as_str(), "string");
908        assert_eq!(ProtoType::Message("User".to_string()).as_str(), "User");
909    }
910}