use prost_reflect::prost::Message;
use prost_reflect::{DescriptorPool, DynamicMessage, MessageDescriptor};
use serde::Serialize;
use crate::httprule::HttpRouter;
pub use crate::httprule::{HttpCall, WsBinding};
#[derive(Debug, thiserror::Error)]
pub enum TranscodeError {
#[error("failed to load descriptor set: {0}")]
Descriptor(String),
#[error("unknown method: {0}")]
UnknownMethod(String),
#[error("json error: {0}")]
Json(#[from] serde_json::Error),
#[error("protobuf decode error: {0}")]
Decode(String),
#[error("http transcoding: {0}")]
Http(String),
}
#[derive(Clone)]
pub struct Transcoder {
pool: DescriptorPool,
router: HttpRouter,
}
impl Transcoder {
pub fn from_file_descriptor_set(bytes: &[u8]) -> Result<Self, TranscodeError> {
let pool = DescriptorPool::decode(bytes)
.map_err(|e| TranscodeError::Descriptor(e.to_string()))?;
let router = HttpRouter::from_pool(&pool);
Ok(Self { pool, router })
}
pub fn has_http_rules(&self) -> bool {
!self.router.is_empty()
}
pub fn transcode_http_request(
&self,
method: &str,
path: &str,
query: Option<&str>,
body: &[u8],
) -> Result<Option<HttpCall>, TranscodeError> {
self.router.transcode(method, path, query, body)
}
pub fn match_ws(&self, path: &str, query: Option<&str>) -> Option<WsBinding> {
self.router.match_ws(path, query)
}
pub fn is_annotated_method(&self, grpc_method: &str) -> bool {
self.router.is_annotated(grpc_method)
}
fn io_types(&self, path: &str) -> Result<(MessageDescriptor, MessageDescriptor), TranscodeError> {
let (service, method) = path
.trim_start_matches('/')
.split_once('/')
.ok_or_else(|| TranscodeError::UnknownMethod(path.to_string()))?;
let svc = self
.pool
.get_service_by_name(service)
.ok_or_else(|| TranscodeError::UnknownMethod(path.to_string()))?;
let m = svc
.methods()
.find(|m| m.name() == method)
.ok_or_else(|| TranscodeError::UnknownMethod(path.to_string()))?;
Ok((m.input(), m.output()))
}
pub fn has_method(&self, path: &str) -> bool {
self.io_types(path).is_ok()
}
pub fn request_json_to_proto(&self, path: &str, json: &[u8]) -> Result<Vec<u8>, TranscodeError> {
let (input, _) = self.io_types(path)?;
Ok(self.json_to_proto(input, json)?.encode_to_vec())
}
pub fn response_proto_to_json(&self, path: &str, proto: &[u8]) -> Result<Vec<u8>, TranscodeError> {
self.response_proto_to_json_body(path, proto, "")
}
pub fn response_proto_to_json_body(
&self,
path: &str,
proto: &[u8],
response_body: &str,
) -> Result<Vec<u8>, TranscodeError> {
let (_, output) = self.io_types(path)?;
if response_body.is_empty() {
return self.proto_to_json(output, proto);
}
let field = crate::httprule::field_by_any_name(&output, response_body)
.ok_or_else(|| TranscodeError::Http(format!("unknown response_body field: {response_body}")))?;
let whole = self.proto_to_json(output, proto)?;
let obj: serde_json::Value = serde_json::from_slice(&whole)?;
let value = obj.get(field.json_name()).cloned().unwrap_or_else(|| json_zero(&field));
Ok(serde_json::to_vec(&value)?)
}
fn json_to_proto(
&self,
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 proto_to_json(&self, desc: MessageDescriptor, proto: &[u8]) -> Result<Vec<u8>, TranscodeError> {
let msg =
DynamicMessage::decode(desc, proto).map_err(|e| TranscodeError::Decode(e.to_string()))?;
let mut buf = Vec::new();
let mut ser = serde_json::Serializer::new(&mut buf);
msg.serialize(&mut ser)?;
Ok(buf)
}
}
fn json_zero(field: &prost_reflect::FieldDescriptor) -> serde_json::Value {
use prost_reflect::Kind;
use serde_json::Value;
if field.is_list() {
return Value::Array(Vec::new());
}
if field.is_map() {
return Value::Object(serde_json::Map::new());
}
match field.kind() {
Kind::Message(_) => Value::Null,
Kind::String => Value::String(String::new()),
Kind::Bytes => Value::String(String::new()),
Kind::Bool => Value::Bool(false),
Kind::Int64 | Kind::Sint64 | Kind::Sfixed64 | Kind::Uint64 | Kind::Fixed64 => {
Value::String("0".to_string())
}
Kind::Enum(e) => match e.get_value(0) {
Some(v) => Value::String(v.name().to_string()),
None => Value::from(0),
},
_ => Value::from(0),
}
}
#[cfg(test)]
mod tests {
use super::*;
use prost_reflect::prost_types::{
field_descriptor_proto::{Label, Type},
DescriptorProto, EnumDescriptorProto, EnumValueDescriptorProto, FieldDescriptorProto,
FileDescriptorProto, FileDescriptorSet,
};
fn zero_fixture() -> MessageDescriptor {
let field = |name: &str, number: i32, ty: Type| FieldDescriptorProto {
name: Some(name.to_string()),
number: Some(number),
label: Some(Label::Optional as i32),
r#type: Some(ty as i32),
..Default::default()
};
let typed = |name: &str, number: i32, ty: Type, type_name: &str| FieldDescriptorProto {
type_name: Some(type_name.to_string()),
..field(name, number, ty)
};
let mut tags = field("tags", 20, Type::String);
tags.label = Some(Label::Repeated as i32);
let file = FileDescriptorProto {
name: Some("zero_fixture.proto".to_string()),
package: Some("webnext.zero.test".to_string()),
syntax: Some("proto3".to_string()),
enum_type: vec![EnumDescriptorProto {
name: Some("Color".to_string()),
value: vec![
EnumValueDescriptorProto {
name: Some("COLOR_UNSPECIFIED".to_string()),
number: Some(0),
..Default::default()
},
EnumValueDescriptorProto {
name: Some("RED".to_string()),
number: Some(1),
..Default::default()
},
],
..Default::default()
}],
message_type: vec![
DescriptorProto {
name: Some("Nested".to_string()),
field: vec![field("id", 1, Type::String)],
..Default::default()
},
DescriptorProto {
name: Some("Req".to_string()),
field: vec![
field("name", 1, Type::String),
field("count", 2, Type::Uint32),
field("big", 3, Type::Int64),
field("ratio", 4, Type::Double),
field("flag", 5, Type::Bool),
field("blob", 6, Type::Bytes),
typed("color", 7, Type::Enum, ".webnext.zero.test.Color"),
typed("nested", 8, Type::Message, ".webnext.zero.test.Nested"),
tags,
],
..Default::default()
},
],
..Default::default()
};
let pool = DescriptorPool::from_file_descriptor_set(FileDescriptorSet { file: vec![file] })
.expect("build fixture pool");
pool.get_message_by_name("webnext.zero.test.Req").expect("Req")
}
#[test]
fn json_zero_per_kind() {
let desc = zero_fixture();
let cases = [
("name", "\"\""), ("count", "0"), ("big", "\"0\""), ("ratio", "0"), ("flag", "false"), ("blob", "\"\""), ("color", "\"COLOR_UNSPECIFIED\""), ("nested", "null"), ("tags", "[]"), ];
for (name, want) in cases {
let field = desc.get_field_by_name(name).expect(name);
let got = serde_json::to_string(&json_zero(&field)).expect("encode");
assert_eq!(got, want, "json_zero({name})");
}
}
#[test]
fn response_body_extraction() {
let tc = Transcoder::from_file_descriptor_set(testecho::FILE_DESCRIPTOR_SET).expect("transcoder");
const METHOD: &str = "/echo.v1.Echo/Unary";
let proto = tc.request_json_to_proto(METHOD, br#"{"message":"hi"}"#).expect("encode");
let body = tc.response_proto_to_json_body(METHOD, &proto, "message").expect("extract");
assert_eq!(String::from_utf8(body).unwrap(), "\"hi\"");
let whole = tc.response_proto_to_json(METHOD, &proto).expect("whole");
assert_eq!(String::from_utf8(whole).unwrap(), r#"{"message":"hi"}"#);
let empty = tc.request_json_to_proto(METHOD, b"{}").expect("encode empty");
let body = tc.response_proto_to_json_body(METHOD, &empty, "message").expect("extract");
assert_eq!(String::from_utf8(body).unwrap(), "\"\"");
assert!(tc.response_proto_to_json_body(METHOD, &proto, "nope").is_err());
}
}