use std::env;
use std::fs;
use std::path::Path;
fn main() {
println!("cargo:rerun-if-changed=proto");
let out_dir = env::var("OUT_DIR").expect("OUT_DIR is set by cargo");
let proto_dir = Path::new(env!("CARGO_MANIFEST_DIR")).join("proto");
let mut generated = String::from(
"// @generated by build.rs from proto/*.proto — do not edit by hand.\n\
// Uses the crate's own protobuf wire codec; no third-party crates.\n\n",
);
if proto_dir.is_dir() {
let mut entries: Vec<_> = fs::read_dir(&proto_dir)
.expect("read proto dir")
.filter_map(|e| e.ok())
.filter(|e| e.path().extension().map(|x| x == "proto").unwrap_or(false))
.collect();
entries.sort_by_key(|e| e.file_name());
let mut compiled: Vec<(String, ProtoFile)> = Vec::new();
for entry in entries {
let file_name = entry.file_name().to_string_lossy().into_owned();
let text = fs::read_to_string(entry.path()).unwrap_or_else(|error| {
panic!("failed to read {}: {error}", entry.path().display())
});
match parse_proto(&text) {
Ok(file) => {
generate_file(&mut generated, &file);
compiled.push((file_name, file));
}
Err(error) => panic!("failed to parse {}: {error}", entry.path().display()),
}
}
generate_descriptors(&mut generated, &compiled);
}
let out_path = Path::new(&out_dir).join("courierust_generated.rs");
fs::write(out_path, &generated).expect("write generated code");
}
fn generate_descriptors(out: &mut String, files: &[(String, ProtoFile)]) {
out.push_str(
"\n/// Protobuf descriptors for the compiled `.proto` files.\n\
///\n\
/// Each constant is a serialized `FileDescriptorProto` — the same bytes\n\
/// `protoc --descriptor_set_out` would write for that file — and the\n\
/// tables beside them are what a server-side reflection implementation\n\
/// answers with. `FILES` handles `file_by_filename`, `SERVICES` backs\n\
/// `list_services`, and a symbol lookup scans the services and the\n\
/// messages in the descriptors.\n\
pub mod descriptors {\n",
);
for (name, file) in files {
let const_name = descriptor_const_name(name);
let bytes = encode_file_descriptor(name, file);
out.push_str(&format!(
" /// `{name}` as a serialized `FileDescriptorProto`.\n pub const {const_name}: &[u8] = &["
));
for (index, byte) in bytes.iter().enumerate() {
if index % 16 == 0 {
out.push_str("\n ");
}
out.push_str(&format!("{byte}, "));
}
out.push_str("\n ];\n");
}
out.push_str(
"\n /// Every compiled `.proto`, by the name it was compiled under.\n \
pub const FILES: &[(&str, &[u8])] = &[\n",
);
for (name, _) in files {
out.push_str(&format!(
" (\"{name}\", {}),\n",
descriptor_const_name(name)
));
}
out.push_str(" ];\n");
out.push_str(
"\n /// Fully-qualified names of every service, in file order.\n \
pub const SERVICES: &[&str] = &[\n",
);
for (_, file) in files {
for service in &file.services {
out.push_str(&format!(
" \"{}\",\n",
qualified(&service.name, &file.package)
));
}
}
out.push_str(" ];\n");
out.push_str(
"\n /// Fully-qualified names of every message, in file order.\n \
pub const MESSAGES: &[&str] = &[\n",
);
for (_, file) in files {
for message in &file.messages {
out.push_str(&format!(
" \"{}\",\n",
qualified(&message.name, &file.package)
));
}
}
out.push_str(" ];\n");
out.push_str(
"\n /// Every symbol a file defines — `package.Service`,\n \
/// `package.Service.Method` and `package.Message` — mapped to the file\n \
/// that defines it, which is what `file_containing_symbol` answers from.\n \
pub const SYMBOLS: &[(&str, &str)] = &[\n",
);
for (name, file) in files {
for service in &file.services {
let service_name = qualified(&service.name, &file.package);
out.push_str(&format!(" (\"{service_name}\", \"{name}\"),\n"));
for rpc in &service.rpcs {
out.push_str(&format!(
" (\"{service_name}.{}\", \"{name}\"),\n",
rpc.name
));
}
}
for message in &file.messages {
out.push_str(&format!(
" (\"{}\", \"{name}\"),\n",
qualified(&message.name, &file.package)
));
}
}
out.push_str(" ];\n}\n");
}
fn descriptor_const_name(file_name: &str) -> String {
let stem = file_name.strip_suffix(".proto").unwrap_or(file_name);
let mut name = String::new();
for ch in stem.chars() {
if ch.is_ascii_alphanumeric() {
name.push(ch.to_ascii_uppercase());
} else {
name.push('_');
}
}
format!("{name}_FILE_DESCRIPTOR")
}
fn qualified(name: &str, package: &Option<String>) -> String {
match package {
Some(package) => format!("{package}.{name}"),
None => name.to_string(),
}
}
fn pb_varint(out: &mut Vec<u8>, mut value: u64) {
while value >= 0x80 {
out.push(value as u8 | 0x80);
value >>= 7;
}
out.push(value as u8);
}
fn pb_tag(out: &mut Vec<u8>, field: u32, wire: u32) {
pb_varint(out, u64::from(field) << 3 | u64::from(wire));
}
fn pb_varint_field(out: &mut Vec<u8>, field: u32, value: u64) {
pb_tag(out, field, 0);
pb_varint(out, value);
}
fn pb_bytes_field(out: &mut Vec<u8>, field: u32, value: &[u8]) {
pb_tag(out, field, 2);
pb_varint(out, value.len() as u64);
out.extend_from_slice(value);
}
fn pb_string_field(out: &mut Vec<u8>, field: u32, value: &str) {
pb_bytes_field(out, field, value.as_bytes());
}
fn pb_message_field(out: &mut Vec<u8>, field: u32, message: &[u8]) {
pb_bytes_field(out, field, message);
}
fn descriptor_type(name: &str) -> u64 {
match name {
"double" => 1,
"float" => 2,
"int64" => 3,
"uint64" => 4,
"int32" => 5,
"fixed64" => 6,
"fixed32" => 7,
"bool" => 8,
"string" => 9,
"bytes" => 12,
"uint32" => 13,
"sfixed32" => 15,
"sfixed64" => 16,
"sint32" => 17,
"sint64" => 18,
other => panic!("no descriptor type for scalar `{other}`"),
}
}
fn type_name(name: &str, package: &Option<String>) -> String {
if name.starts_with('.') {
return name.to_string();
}
if name.contains('.') {
return format!(".{name}");
}
match package {
Some(package) => format!(".{package}.{name}"),
None => format!(".{name}"),
}
}
fn json_name(field: &str) -> String {
let mut out = String::with_capacity(field.len());
let mut upper_next = false;
for ch in field.chars() {
if ch == '_' {
upper_next = true;
continue;
}
if upper_next {
out.extend(ch.to_uppercase());
upper_next = false;
} else {
out.push(ch);
}
}
out
}
fn encode_file_descriptor(file_name: &str, file: &ProtoFile) -> Vec<u8> {
let mut out = Vec::new();
pb_string_field(&mut out, 1, file_name); if let Some(package) = &file.package {
pb_string_field(&mut out, 2, package); }
for message in &file.messages {
pb_message_field(&mut out, 4, &encode_message(message, &file.package));
}
for service in &file.services {
pb_message_field(&mut out, 6, &encode_service(service, &file.package));
}
pb_string_field(&mut out, 12, "proto3"); out
}
fn encode_message(message: &Message, package: &Option<String>) -> Vec<u8> {
let mut out = Vec::new();
pb_string_field(&mut out, 1, &message.name);
for field in &message.fields {
let (label, ty, type_name) = match &field.ty {
FieldType::Scalar(scalar) => (1, descriptor_type(scalar), None),
FieldType::Repeated(scalar) => (3, descriptor_type(scalar), None),
FieldType::Message(name) => (1, 11, Some(type_name(name, package))),
FieldType::RepeatedMessage(name) => (3, 11, Some(type_name(name, package))),
};
let mut encoded = Vec::new();
pb_string_field(&mut encoded, 1, &field.name);
pb_varint_field(&mut encoded, 3, u64::from(field.number));
pb_varint_field(&mut encoded, 4, label);
pb_varint_field(&mut encoded, 5, ty);
if let Some(type_name) = type_name {
pb_string_field(&mut encoded, 6, &type_name);
}
pb_string_field(&mut encoded, 10, &json_name(&field.name));
pb_message_field(&mut out, 2, &encoded);
}
out
}
fn encode_service(service: &Service, package: &Option<String>) -> Vec<u8> {
let mut out = Vec::new();
pb_string_field(&mut out, 1, &service.name);
for rpc in &service.rpcs {
let mut encoded = Vec::new();
pb_string_field(&mut encoded, 1, &rpc.name);
pb_string_field(&mut encoded, 2, &type_name(&rpc.request, package));
pb_string_field(&mut encoded, 3, &type_name(&rpc.response, package));
if rpc.server_streaming {
pb_varint_field(&mut encoded, 6, 1);
}
pb_message_field(&mut out, 2, &encoded);
}
out
}
#[derive(Debug, Clone, PartialEq)]
enum FieldType {
Scalar(&'static str),
Repeated(&'static str),
Message(String),
RepeatedMessage(String),
}
#[derive(Debug)]
struct Field {
number: u32,
name: String,
ty: FieldType,
}
#[derive(Debug)]
struct Message {
name: String,
fields: Vec<Field>,
}
#[derive(Debug)]
struct Rpc {
name: String,
request: String,
response: String,
server_streaming: bool,
}
#[derive(Debug)]
struct Service {
name: String,
rpcs: Vec<Rpc>,
}
#[derive(Debug, Default)]
struct ProtoFile {
package: Option<String>,
messages: Vec<Message>,
services: Vec<Service>,
}
fn parse_proto(text: &str) -> Result<ProtoFile, String> {
let text = text.strip_prefix('\u{feff}').unwrap_or(text);
let mut tokens = tokenize(text)?;
let mut file = ProtoFile::default();
while !tokens.is_empty() && !matches!(tokens.first(), Some(Token::Eof)) {
let (keyword, semicolon) = match tokens.remove(0) {
Token::Ident(s) => (s, false),
Token::Semicolon => (String::new(), true),
other => return Err(format!("expected identifier, got {other:?}")),
};
if semicolon {
continue;
}
match keyword.as_str() {
"syntax" => {
expect(&mut tokens, Token::Eq, "syntax =")?;
tokens.remove(0); expect(&mut tokens, Token::Semicolon, "syntax ;")?;
}
"package" => {
let mut parts = Vec::new();
loop {
match tokens.remove(0) {
Token::Ident(s) => parts.push(s),
Token::Dot => {}
Token::Semicolon => break,
other => return Err(format!("bad package: {other:?}")),
}
}
file.package = Some(parts.join("."));
}
"message" => {
let name = expect_ident(&mut tokens, "message name")?;
expect(&mut tokens, Token::LBrace, "message {")?;
let mut fields = Vec::new();
loop {
match tokens.remove(0) {
Token::RBrace => break,
Token::Ident(leading) => {
let mut ty_name = leading;
let mut repeated = false;
if ty_name == "repeated" {
repeated = true;
ty_name = expect_ident(&mut tokens, "repeated type")?;
}
let field_name = expect_ident(&mut tokens, "field name")?;
expect(&mut tokens, Token::Eq, "field =")?;
let number = expect_u32(&mut tokens, "field number")?;
if number == 0 || number >= (1 << 29) {
return Err(format!("invalid field number {number} in {name}"));
}
if fields.iter().any(|f: &Field| f.number == number) {
return Err(format!(
"duplicate field number {number} in message {name}"
));
}
expect(&mut tokens, Token::Semicolon, "field ;")?;
let ty = if let Some(scalar) = scalar_of(&ty_name) {
if repeated {
FieldType::Repeated(scalar)
} else {
FieldType::Scalar(scalar)
}
} else if repeated {
FieldType::RepeatedMessage(ty_name)
} else {
FieldType::Message(ty_name)
};
fields.push(Field {
number,
name: field_name,
ty,
});
}
other => return Err(format!("bad field: {other:?}")),
}
}
file.messages.push(Message { name, fields });
}
"service" => {
let name = expect_ident(&mut tokens, "service name")?;
expect(&mut tokens, Token::LBrace, "service {")?;
let mut rpcs = Vec::new();
loop {
match tokens.remove(0) {
Token::RBrace => break,
Token::Ident(word) if word == "rpc" => {
let rpc_name = expect_ident(&mut tokens, "rpc name")?;
expect(&mut tokens, Token::LParen, "rpc (")?;
let request = expect_ident(&mut tokens, "request type")?;
expect(&mut tokens, Token::RParen, "rpc )")?;
expect(&mut tokens, Token::Ident("returns".to_owned()), "returns")?;
expect(&mut tokens, Token::LParen, "returns (")?;
let mut server_streaming = false;
if matches!(tokens.first(), Some(Token::Ident(w)) if w == "stream") {
tokens.remove(0);
server_streaming = true;
}
let response = expect_ident(&mut tokens, "response type")?;
expect(&mut tokens, Token::RParen, "returns )")?;
expect(&mut tokens, Token::Semicolon, "rpc ;")?;
rpcs.push(Rpc {
name: rpc_name,
request,
response,
server_streaming,
});
}
other => return Err(format!("bad service item: {other:?}")),
}
}
file.services.push(Service { name, rpcs });
}
"import" => {
let target = match tokens.remove(0) {
Token::Quoted(name) => name,
other => return Err(format!("unsupported import form: {other:?}")),
};
match tokens.remove(0) {
Token::Semicolon => {}
Token::Eof => return Err(format!("unterminated import of {target}")),
other => return Err(format!("unexpected {other:?} in import of {target}")),
}
return Err(format!(
"`import {target}` is not supported: every .proto compiles standalone"
));
}
"option" => {
loop {
match tokens.remove(0) {
Token::Semicolon => break,
Token::Eof => return Err("unterminated option".to_string()),
_ => {}
}
}
}
other => return Err(format!("unsupported proto keyword: {other}")),
}
}
for message in &file.messages {
for field in &message.fields {
let referenced = match &field.ty {
FieldType::Message(name) | FieldType::RepeatedMessage(name) => name,
_ => continue,
};
if !file.messages.iter().any(|m| &m.name == referenced) {
return Err(format!(
"message {}: field {} has type `{}`, which is neither an unsupported-by-design \
scalar (`sfixed32`/`sfixed64`) nor a message declared in this file",
message.name, field.name, referenced
));
}
}
}
Ok(file)
}
fn scalar_of(name: &str) -> Option<&'static str> {
Some(match name {
"string" => "string",
"bytes" => "bytes",
"bool" => "bool",
"int32" => "int32",
"int64" => "int64",
"uint32" => "uint32",
"uint64" => "uint64",
"sint32" => "sint32",
"sint64" => "sint64",
"fixed32" => "fixed32",
"fixed64" => "fixed64",
"float" => "float",
"double" => "double",
_ => return None,
})
}
#[derive(Debug, Clone, PartialEq)]
enum Token {
Ident(String),
Eq,
Dot,
Semicolon,
LBrace,
RBrace,
LParen,
RParen,
Quoted(String),
Number(u32),
Eof,
}
fn tokenize(text: &str) -> Result<Vec<Token>, String> {
let bytes = text.as_bytes();
let mut tokens = Vec::new();
let mut i = 0usize;
while i < bytes.len() {
let c = bytes[i] as char;
if c.is_whitespace() {
i += 1;
continue;
}
if c == '/' && i + 1 < bytes.len() && bytes[i + 1] as char == '/' {
while i < bytes.len() && bytes[i] as char != '\n' {
i += 1;
}
continue;
}
if c == '/' && i + 1 < bytes.len() && bytes[i + 1] as char == '*' {
let end = text[i + 2..]
.find("*/")
.ok_or_else(|| "unterminated block comment".to_string())?;
i += 2 + end + 2;
continue;
}
match c {
'=' => {
tokens.push(Token::Eq);
i += 1;
}
'.' => {
tokens.push(Token::Dot);
i += 1;
}
';' => {
tokens.push(Token::Semicolon);
i += 1;
}
'{' => {
tokens.push(Token::LBrace);
i += 1;
}
'}' => {
tokens.push(Token::RBrace);
i += 1;
}
'(' => {
tokens.push(Token::LParen);
i += 1;
}
')' => {
tokens.push(Token::RParen);
i += 1;
}
'"' | '\'' => {
let quote = c;
let start = i + 1;
let mut end = start;
while end < bytes.len() && bytes[end] as char != quote {
end += 1;
}
if end >= bytes.len() {
return Err("unterminated string literal".to_string());
}
tokens.push(Token::Quoted(text[start..end].to_string()));
i = end + 1;
}
'0'..='9' => {
let start = i;
while i < bytes.len() && (bytes[i] as char).is_ascii_digit() {
i += 1;
}
let value: u32 = text[start..i]
.parse()
.map_err(|_| "invalid field number".to_string())?;
tokens.push(Token::Number(value));
}
c if c.is_ascii_alphabetic() || c == '_' => {
let start = i;
while i < bytes.len()
&& (bytes[i] as char == '_' || (bytes[i] as char).is_ascii_alphanumeric())
{
i += 1;
}
tokens.push(Token::Ident(text[start..i].to_string()));
}
other if !other.is_ascii() => {
return Err(format!(
"non-ASCII character at offset {i} (0x{:02x}): proto3 source is ASCII \
outside comments and string literals",
bytes[i]
));
}
other => return Err(format!("unexpected character `{other}`")),
}
}
tokens.push(Token::Eof);
Ok(tokens)
}
fn expect(tokens: &mut Vec<Token>, expected: Token, what: &str) -> Result<(), String> {
if tokens.is_empty() {
return Err(format!("unexpected end of file while parsing {what}"));
}
let actual = tokens.remove(0);
if actual == expected {
Ok(())
} else {
Err(format!("expected {what}, got {actual:?}"))
}
}
fn expect_ident(tokens: &mut Vec<Token>, what: &str) -> Result<String, String> {
match tokens.remove(0) {
Token::Ident(s) => Ok(s),
other => Err(format!("expected {what}, got {other:?}")),
}
}
fn expect_u32(tokens: &mut Vec<Token>, what: &str) -> Result<u32, String> {
match tokens.remove(0) {
Token::Number(n) => Ok(n),
other => Err(format!("expected {what}, got {other:?}")),
}
}
fn to_snake(name: &str) -> String {
let mut out = String::new();
for (i, c) in name.chars().enumerate() {
if c.is_uppercase() {
if i > 0 {
out.push('_');
}
out.push(c.to_ascii_lowercase());
} else {
out.push(c);
}
}
out
}
fn rust_scalar(name: &str) -> &'static str {
match name {
"string" => "alloc::string::String",
"bytes" => "alloc::vec::Vec<u8>",
"bool" => "bool",
"int32" | "sint32" => "i32",
"int64" | "sint64" => "i64",
"uint32" | "fixed32" => "u32",
"uint64" | "fixed64" => "u64",
"float" => "f32",
"double" => "f64",
_ => unreachable!("scalar"),
}
}
fn rust_scalar_f32_f64(name: &str) -> &'static str {
rust_scalar(name)
}
fn generate_file(out: &mut String, file: &ProtoFile) {
if let Some(package) = &file.package {
let parts: Vec<&str> = package.split('.').collect();
let mut full_pkg = String::new();
for (i, part) in parts.iter().enumerate() {
if i > 0 {
full_pkg.push('.');
}
full_pkg.push_str(part);
}
for part in &parts {
out.push_str(&format!(
"/// Proto package `{full_pkg}` (generated).\npub mod {part} {{\n"
));
}
out.push_str(" use crate::courierust_grpc::proto;\n");
out.push_str(" use crate::courierust_grpc::proto::ProtoMessage;\n");
out.push_str(" use crate::courierust_grpc::codec;\n");
out.push_str(" use crate::courierust_error::{Error, Result};\n");
for message in &file.messages {
generate_message(out, message, 1);
}
for service in &file.services {
generate_service(out, service, 1, file.package.as_deref());
}
for _ in &parts {
out.push_str("}\n");
}
} else {
out.push_str("use crate::courierust_grpc::proto;\n");
out.push_str("use crate::courierust_grpc::proto::ProtoMessage;\n");
out.push_str("use crate::courierust_grpc::codec;\n");
out.push_str("use crate::courierust_error::{Error, Result};\n");
for message in &file.messages {
generate_message(out, message, 0);
}
for service in &file.services {
generate_service(out, service, 0, file.package.as_deref());
}
}
}
fn indent(level: usize) -> String {
" ".repeat(level)
}
fn scalar_kind(name: &str) -> &'static str {
match name {
"bool" => "bool",
"int32" => "i32",
"int64" => "i64",
"uint32" => "u32",
"uint64" => "u64",
"sint32" => "s32",
"sint64" => "s64",
"fixed32" => "fixed32",
"fixed64" => "fixed64",
"float" => "float",
"double" => "double",
_ => unreachable!("{name} is not a scalar"),
}
}
fn is_length_delimited(name: &str) -> bool {
name == "string" || name == "bytes"
}
fn is_present(name: &str, field: &str) -> String {
match name {
"string" | "bytes" => format!("!self.{field}.is_empty()"),
"bool" => format!("self.{field}"),
"float" | "double" => format!("self.{field}.to_bits() != 0"),
_ => format!("self.{field} != 0"),
}
}
fn generate_message(out: &mut String, message: &Message, level: usize) {
let ind = indent(level);
out.push_str(&format!(
"\n{ind}/// Message `{}` (proto3).\n",
message.name
));
out.push_str(&format!(
"{ind}#[derive(Debug, Clone, PartialEq, Default)]\n"
));
out.push_str(&format!("{ind}pub struct {} {{\n", message.name));
for field in &message.fields {
let rust_ty = match &field.ty {
FieldType::Scalar(s) => rust_scalar_f32_f64(s).to_string(),
FieldType::Repeated(s) => {
format!("alloc::vec::Vec<{}>", rust_scalar_f32_f64(s))
}
FieldType::Message(m) => format!("::core::option::Option<{m}>"),
FieldType::RepeatedMessage(m) => format!("alloc::vec::Vec<{m}>"),
};
out.push_str(&format!(
"{ind} /// Field {} (`{}`).\n",
field.number, field.name
));
out.push_str(&format!("{ind} pub {}: {rust_ty},\n", field.name));
}
out.push_str(&format!("{ind}}}\n"));
out.push_str(&format!(
"\n{ind}impl proto::ProtoMessage for {} {{\n",
message.name
));
out.push_str(&format!(
"{ind} fn encode_message_body(&self, out: &mut alloc::vec::Vec<u8>) {{\n"
));
for field in &message.fields {
let e = indent(level + 2);
let n = field.number;
match &field.ty {
FieldType::Scalar(s) => {
let expr = if is_length_delimited(s) {
format!("&self.{}", field.name)
} else {
format!("self.{}", field.name)
};
out.push_str(&format!(
"{e}if {} {{\n{e} proto::Encoder::{s}(out, {n}, {expr});\n{e}}}\n",
is_present(s, &field.name)
));
}
FieldType::Message(m) => out.push_str(&format!(
"{e}if let Some(value) = &self.{} {{ proto::Encoder::message::<{m}>(out, {n}, value); }}\n",
field.name
)),
FieldType::Repeated(s) if is_length_delimited(s) => out.push_str(&format!(
"{e}for value in &self.{} {{ proto::Encoder::{s}(out, {n}, value); }}\n",
field.name
)),
FieldType::Repeated(s) => {
let kind = scalar_kind(s);
out.push_str(&format!(
"{e}if !self.{}.is_empty() {{\n\
{e} let mut payload = alloc::vec::Vec::new();\n\
{e} for value in &self.{} {{ proto::push_packed_{kind}(&mut payload, *value); }}\n\
{e} proto::encode_packed_payload(out, {n}, &payload);\n\
{e}}}\n",
field.name, field.name
));
}
FieldType::RepeatedMessage(m) => out.push_str(&format!(
"{e}for value in &self.{} {{ proto::Encoder::message::<{m}>(out, {n}, value); }}\n",
field.name
)),
}
}
out.push_str(&format!("{ind} }}\n"));
out.push_str(&format!(
"{ind} fn decode_message_body(buf: &[u8]) -> Result<Self> {{\n"
));
out.push_str(&format!("{ind} let mut result = Self::default();\n"));
out.push_str(&format!(
"{ind} let mut reader = proto::SliceReader::new(buf);\n"
));
out.push_str(&format!("{ind} while reader.remaining() > 0 {{\n"));
out.push_str(&format!(
"{ind} proto::Decoder::field(&mut reader, |number, wire, value, payload| {{\n"
));
out.push_str(&format!("{ind} let _ = value;\n"));
out.push_str(&format!("{ind} match (number, wire) {{\n"));
for field in &message.fields {
let e = indent(level + 4);
let n = field.number;
let f = &field.name;
match &field.ty {
FieldType::Scalar(s) => {
if is_length_delimited(s) {
if *s == "string" {
out.push_str(&format!(
"{e}({n}, proto::WireType::LengthDelimited) => {{\n\
{e} result.{f} = ::core::str::from_utf8(payload)\n\
{e} .map_err(|_| Error::protocol(\"invalid utf8 in protobuf string\"))?\n\
{e} .to_string();\n\
{e}}}\n"
));
} else {
out.push_str(&format!(
"{e}({n}, proto::WireType::LengthDelimited) => {{\n\
{e} result.{f} = payload.to_vec();\n\
{e}}}\n"
));
}
} else {
let wire = match *s {
"fixed32" | "float" => "proto::WireType::Fixed32",
"fixed64" | "double" => "proto::WireType::Fixed64",
_ => "proto::WireType::Varint",
};
let kind = scalar_kind(s);
out.push_str(&format!(
"{e}({n}, {wire}) => {{\n\
{e} result.{f} = proto::conv_{kind}(value);\n\
{e}}}\n"
));
}
}
FieldType::Message(m) => out.push_str(&format!(
"{e}({n}, proto::WireType::LengthDelimited) => {{\n\
{e} result.{f} = Some(proto::decode_sub_message::<{m}>(payload)?);\n\
{e}}}\n"
)),
FieldType::Repeated(s) if *s == "string" => out.push_str(&format!(
"{e}({n}, proto::WireType::LengthDelimited) => {{\n\
{e} result.{f}.push(\n\
{e} ::core::str::from_utf8(payload)\n\
{e} .map_err(|_| Error::protocol(\"invalid utf8 in protobuf string\"))?\n\
{e} .to_string(),\n\
{e} );\n\
{e}}}\n"
)),
FieldType::Repeated(s) if *s == "bytes" => out.push_str(&format!(
"{e}({n}, proto::WireType::LengthDelimited) => {{\n\
{e} result.{f}.push(payload.to_vec());\n\
{e}}}\n"
)),
FieldType::Repeated(s) => {
let kind = scalar_kind(s);
let element_wire = match *s {
"fixed32" | "float" => "proto::WireType::Fixed32",
"fixed64" | "double" => "proto::WireType::Fixed64",
_ => "proto::WireType::Varint",
};
out.push_str(&format!(
"{e}({n}, proto::WireType::LengthDelimited) => {{\n\
{e} for v in proto::decode_packed_{kind}(payload)? {{\n\
{e} result.{f}.push(v);\n\
{e} }}\n\
{e}}}\n\
{e}({n}, {element_wire}) => {{\n\
{e} result.{f}.push(proto::conv_{kind}(value));\n\
{e}}}\n"
));
}
FieldType::RepeatedMessage(m) => out.push_str(&format!(
"{e}({n}, proto::WireType::LengthDelimited) => {{\n\
{e} result.{f}.push(proto::decode_sub_message::<{m}>(payload)?);\n\
{e}}}\n"
)),
}
}
out.push_str(&format!("{ind} _ => {{}}\n"));
out.push_str(&format!("{ind} }}\n"));
out.push_str(&format!("{ind} Ok(())\n"));
out.push_str(&format!("{ind} }})?;\n"));
out.push_str(&format!("{ind} }}\n"));
out.push_str(&format!("{ind} Ok(result)\n"));
out.push_str(&format!("{ind} }}\n"));
out.push_str(&format!("{ind}}}\n"));
out.push_str(&format!(
"\n{ind}impl codec::EncodeMessage for {} {{\n",
message.name
));
out.push_str(&format!(
"{ind} fn encode_message(&self) -> Result<alloc::vec::Vec<u8>> {{\n\
{ind} let mut out = alloc::vec::Vec::new();\n\
{ind} self.encode_message_body(&mut out);\n\
{ind} Ok(out)\n\
{ind} }}\n\
{ind}}}\n"
));
out.push_str(&format!(
"\n{ind}impl codec::DecodeMessage for {} {{\n",
message.name
));
out.push_str(&format!(
"{ind} fn decode_message(bytes: &[u8]) -> Result<Self> {{\n\
{ind} Self::decode_message_body(bytes)\n\
{ind} }}\n\
{ind}}}\n"
));
}
fn generate_service(out: &mut String, service: &Service, level: usize, package: Option<&str>) {
let ind = indent(level);
let qualified = match package {
Some(pkg) => format!("/{pkg}.{}", service.name),
None => format!("/{}", service.name),
};
let path = qualified.clone();
out.push_str(&format!(
"\n{ind}/// Generated gRPC client stub for service `{}`.\n",
service.name
));
out.push_str(&format!(
"{ind}#[derive(Clone)]\n{ind}pub struct {} {{\n{ind} inner: crate::courierust_grpc::GrpcClient,\n{ind}}}\n",
service.name
));
out.push_str(&format!("\n{ind}impl {} {{\n", service.name));
out.push_str(&format!(
"{ind} /// Create a typed client stub bound to a gRPC connection.\n\
{ind} pub fn new(client: crate::courierust_grpc::GrpcClient) -> Self {{\n\
{ind} Self {{ inner: client }}\n\
{ind} }}\n"
));
out.push_str(&format!(
"{ind} /// The fully-qualified service name.\n\
{ind} pub fn service_path(&self) -> &'static str {{\n\
{ind} \"{path}\"\n\
{ind} }}\n"
));
for rpc in &service.rpcs {
let method = to_snake(&rpc.name);
let full = format!("{path}/{}", rpc.name);
if rpc.server_streaming {
out.push_str(&format!(
"\n{ind} /// Server-streaming RPC `{full}`: returns a stream of `{}` messages.\n\
{ind} pub fn {method}(&self, request: {}) -> Result<crate::courierust_grpc::MessageStream> {{\n\
{ind} let bytes = codec::EncodeMessage::encode_message(&request)?;\n\
{ind} self.inner.call_stream(\"{full}\", bytes.into())\n\
{ind} }}\n",
rpc.response, rpc.request
));
} else {
out.push_str(&format!(
"\n{ind} /// Unary RPC `{full}`.\n\
{ind} pub fn {method}(&self, request: {}) -> Result<{}> {{\n\
{ind} self.inner.call_unary(\"{full}\", &request)\n\
{ind} }}\n",
rpc.request, rpc.response
));
}
}
out.push_str(&format!("{ind}}}\n"));
}