use prost_reflect::prost::Message;
use prost_reflect::{DescriptorPool, DynamicMessage, Kind, MessageDescriptor, ReflectMessage, Value};
use crate::transcode::TranscodeError;
const HTTP_EXT: &str = "google.api.http";
#[derive(Debug, Clone)]
enum Segment {
Literal(String),
Single(Vec<String>),
Rest(Vec<String>),
}
#[derive(Debug, Clone)]
enum BodyRule {
Wildcard,
Field(String),
None,
}
#[derive(Debug, Clone)]
struct Binding {
http_method: String, segments: Vec<Segment>,
body: BodyRule,
grpc_method: String, input: MessageDescriptor,
}
type PathVar = (Vec<String>, String);
pub struct HttpCall {
pub grpc_method: String,
pub message: Vec<u8>,
}
pub struct WsBinding {
binding: Binding,
vars: Vec<PathVar>,
query: Option<String>,
}
impl WsBinding {
pub fn grpc_method(&self) -> &str {
&self.binding.grpc_method
}
pub fn has_body(&self) -> bool {
!matches!(self.binding.body, BodyRule::None)
}
pub fn build_message(&self, body: &[u8]) -> Result<Vec<u8>, TranscodeError> {
build_message(&self.binding, &self.vars, self.query.as_deref(), body)
}
}
#[derive(Clone, Default)]
pub struct HttpRouter {
bindings: Vec<Binding>,
}
impl HttpRouter {
pub fn from_pool(pool: &DescriptorPool) -> Self {
let mut bindings = Vec::new();
let Some(ext) = pool.get_extension_by_name(HTTP_EXT) else {
return Self { bindings };
};
for service in pool.services() {
for method in service.methods() {
let opts = method.options();
if !opts.has_extension(&ext) {
continue;
}
let grpc_method = format!("/{}/{}", service.full_name(), method.name());
let input = method.input();
let rule = opts.get_extension(&ext);
if let Some(rule) = rule.as_message() {
collect_rule(&mut bindings, rule, &grpc_method, &input);
}
}
}
Self { bindings }
}
pub fn is_empty(&self) -> bool {
self.bindings.is_empty()
}
pub fn match_ws(&self, path: &str, query: Option<&str>) -> Option<WsBinding> {
for b in &self.bindings {
if let Some(vars) = match_segments(&b.segments, path) {
return Some(WsBinding {
binding: b.clone(),
vars,
query: query.map(str::to_string),
});
}
}
None
}
pub fn is_annotated(&self, grpc_method: &str) -> bool {
self.bindings.iter().any(|b| b.grpc_method == grpc_method)
}
fn match_request(&self, method: &str, path: &str) -> Option<(&Binding, Vec<PathVar>)> {
let want = method.to_ascii_uppercase();
for b in &self.bindings {
if b.http_method != want {
continue;
}
if let Some(vars) = match_segments(&b.segments, path) {
return Some((b, vars));
}
}
None
}
pub fn transcode(
&self,
method: &str,
path: &str,
query: Option<&str>,
body: &[u8],
) -> Result<Option<HttpCall>, TranscodeError> {
let Some((binding, vars)) = self.match_request(method, path) else {
return Ok(None);
};
let message = build_message(binding, &vars, query, body)?;
Ok(Some(HttpCall { grpc_method: binding.grpc_method.clone(), message }))
}
}
fn collect_rule(bindings: &mut Vec<Binding>, rule: &DynamicMessage, grpc_method: &str, input: &MessageDescriptor) {
push_binding(bindings, rule, grpc_method, input);
if let Some(list) = rule.get_field_by_name("additional_bindings") {
if let Some(items) = list.as_list() {
for item in items {
if let Some(m) = item.as_message() {
push_binding(bindings, m, grpc_method, input);
}
}
}
}
}
fn push_binding(bindings: &mut Vec<Binding>, rule: &DynamicMessage, grpc_method: &str, input: &MessageDescriptor) {
let Some((http_method, template)) = verb_and_path(rule) else {
return;
};
bindings.push(Binding {
http_method,
segments: parse_template(&template),
body: body_rule(rule),
grpc_method: grpc_method.to_string(),
input: input.clone(),
});
}
fn verb_and_path(rule: &DynamicMessage) -> Option<(String, String)> {
for (field, verb) in [("get", "GET"), ("put", "PUT"), ("post", "POST"), ("delete", "DELETE"), ("patch", "PATCH")] {
if let Some(v) = rule.get_field_by_name(field) {
if let Some(s) = v.as_str() {
if !s.is_empty() {
return Some((verb.to_string(), s.to_string()));
}
}
}
}
if let Some(v) = rule.get_field_by_name("custom") {
if let Some(m) = v.as_message() {
let kind = m.get_field_by_name("kind").and_then(|k| k.as_str().map(str::to_string)).unwrap_or_default();
let path = m.get_field_by_name("path").and_then(|p| p.as_str().map(str::to_string)).unwrap_or_default();
if !kind.is_empty() {
return Some((kind.to_ascii_uppercase(), path));
}
}
}
None
}
fn body_rule(rule: &DynamicMessage) -> BodyRule {
match rule.get_field_by_name("body").and_then(|b| b.as_str().map(str::to_string)).unwrap_or_default().as_str() {
"" => BodyRule::None,
"*" => BodyRule::Wildcard,
field => BodyRule::Field(field.to_string()),
}
}
fn parse_template(template: &str) -> Vec<Segment> {
let path = template.split(':').next().unwrap_or(template);
path.trim_matches('/')
.split('/')
.filter(|s| !s.is_empty())
.map(|seg| {
if let Some(inner) = seg.strip_prefix('{').and_then(|s| s.strip_suffix('}')) {
let (field, pattern) = inner.split_once('=').unwrap_or((inner, "*"));
let field_path: Vec<String> = field.split('.').map(str::to_string).collect();
if pattern == "**" {
Segment::Rest(field_path)
} else {
Segment::Single(field_path)
}
} else {
Segment::Literal(seg.to_string())
}
})
.collect()
}
fn match_segments(segments: &[Segment], path: &str) -> Option<Vec<PathVar>> {
let parts: Vec<&str> = path.trim_matches('/').split('/').filter(|s| !s.is_empty()).collect();
let mut vars = Vec::new();
let mut i = 0;
for seg in segments {
match seg {
Segment::Literal(lit) => {
if parts.get(i) != Some(&lit.as_str()) {
return None;
}
i += 1;
}
Segment::Single(field) => {
let part = parts.get(i)?;
vars.push((field.clone(), percent_decode(part)));
i += 1;
}
Segment::Rest(field) => {
let rest: Vec<String> = parts[i..].iter().map(|p| percent_decode(p)).collect();
vars.push((field.clone(), rest.join("/")));
i = parts.len();
}
}
}
if i == parts.len() {
Some(vars)
} else {
None
}
}
fn build_message(
binding: &Binding,
vars: &[PathVar],
query: Option<&str>,
body: &[u8],
) -> Result<Vec<u8>, TranscodeError> {
let mut msg = match &binding.body {
BodyRule::Wildcard => deserialize_message(binding.input.clone(), body)?,
BodyRule::None => DynamicMessage::new(binding.input.clone()),
BodyRule::Field(field) => {
let mut m = DynamicMessage::new(binding.input.clone());
if !body.is_empty() {
set_message_field(&mut m, field, body)?;
}
m
}
};
for (field_path, value) in vars {
set_by_path(&mut msg, field_path, value)?;
}
if !matches!(binding.body, BodyRule::Wildcard) {
if let Some(q) = query {
for (key, value) in parse_query(q) {
let field_path: Vec<String> = key.split('.').map(str::to_string).collect();
if vars.iter().any(|(fp, _)| *fp == field_path) {
continue; }
set_by_path(&mut msg, &field_path, &value)?;
}
}
}
Ok(msg.encode_to_vec())
}
fn deserialize_message(desc: MessageDescriptor, json: &[u8]) -> Result<DynamicMessage, TranscodeError> {
if json.is_empty() {
return Ok(DynamicMessage::new(desc));
}
let mut de = serde_json::Deserializer::from_slice(json);
let msg = DynamicMessage::deserialize(desc, &mut de)?;
de.end()?;
Ok(msg)
}
fn set_message_field(msg: &mut DynamicMessage, name: &str, json: &[u8]) -> Result<(), TranscodeError> {
let field = msg
.descriptor()
.get_field_by_name(name)
.ok_or_else(|| TranscodeError::Http(format!("unknown body field: {name}")))?;
match field.kind() {
Kind::Message(md) => {
let sub = deserialize_message(md, json)?;
msg.set_field(&field, Value::Message(sub));
Ok(())
}
_ => Err(TranscodeError::Http(format!("body field {name} must be a message"))),
}
}
fn set_by_path(msg: &mut DynamicMessage, path: &[String], raw: &str) -> Result<(), TranscodeError> {
let field = msg
.descriptor()
.get_field_by_name(&path[0])
.ok_or_else(|| TranscodeError::Http(format!("unknown field: {}", path[0])))?;
if path.len() == 1 {
let value = coerce(&field, raw)?;
if field.is_list() {
if let Some(list) = msg.get_field_mut(&field).as_list_mut() {
list.push(value);
}
} else {
msg.set_field(&field, value);
}
Ok(())
} else {
let sub = msg
.get_field_mut(&field)
.as_message_mut()
.ok_or_else(|| TranscodeError::Http(format!("field {} is not a message", path[0])))?;
set_by_path(sub, &path[1..], raw)
}
}
fn coerce(field: &prost_reflect::FieldDescriptor, raw: &str) -> Result<Value, TranscodeError> {
let num = |ok: Option<Value>| ok.ok_or_else(|| TranscodeError::Http(format!("invalid value for {}: {raw:?}", field.name())));
Ok(match field.kind() {
Kind::String => Value::String(raw.to_string()),
Kind::Bool => match raw {
"true" | "1" => Value::Bool(true),
"false" | "0" => Value::Bool(false),
_ => return Err(TranscodeError::Http(format!("invalid bool: {raw:?}"))),
},
Kind::Int32 | Kind::Sint32 | Kind::Sfixed32 => num(raw.parse().ok().map(Value::I32))?,
Kind::Int64 | Kind::Sint64 | Kind::Sfixed64 => num(raw.parse().ok().map(Value::I64))?,
Kind::Uint32 | Kind::Fixed32 => num(raw.parse().ok().map(Value::U32))?,
Kind::Uint64 | Kind::Fixed64 => num(raw.parse().ok().map(Value::U64))?,
Kind::Float => num(raw.parse().ok().map(Value::F32))?,
Kind::Double => num(raw.parse().ok().map(Value::F64))?,
Kind::Enum(e) => {
if let Ok(n) = raw.parse::<i32>() {
Value::EnumNumber(n)
} else {
let v = e
.values()
.find(|v| v.name() == raw)
.ok_or_else(|| TranscodeError::Http(format!("unknown enum value: {raw:?}")))?;
Value::EnumNumber(v.number())
}
}
Kind::Bytes => return Err(TranscodeError::Http("bytes fields cannot bind from path/query".into())),
Kind::Message(_) => return Err(TranscodeError::Http("cannot bind a scalar to a message field".into())),
})
}
fn parse_query(query: &str) -> Vec<(String, String)> {
query
.split('&')
.filter(|p| !p.is_empty())
.map(|pair| {
let (k, v) = pair.split_once('=').unwrap_or((pair, ""));
(decode_query(k), decode_query(v))
})
.collect()
}
fn decode_query(s: &str) -> String {
percent_decode(&s.replace('+', " "))
}
fn percent_decode(s: &str) -> String {
let bytes = s.as_bytes();
let mut out = Vec::with_capacity(bytes.len());
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'%' && i + 2 < bytes.len() {
if let Ok(b) = u8::from_str_radix(&s[i + 1..i + 3], 16) {
out.push(b);
i += 3;
continue;
}
}
out.push(bytes[i]);
i += 1;
}
String::from_utf8_lossy(&out).into_owned()
}