use super::*;
use serde::de::{Deserializer, Visitor};
use serde_json::value::RawValue;
use std::{fmt, ops::Range};
#[derive(Debug)]
struct Text {
bytes: [u8; MAX_ID_BYTES],
len: usize,
oversized: bool,
}
impl Default for Text {
fn default() -> Self {
Self {
bytes: [0; MAX_ID_BYTES],
len: 0,
oversized: false,
}
}
}
impl Text {
fn get(&self) -> &str {
std::str::from_utf8(&self.bytes[..self.len]).expect("copied whole UTF-8")
}
fn identity(&self) -> Result<(), CodecError> {
validate_id(self.get())
}
fn trace(&self) -> Result<(), CodecError> {
if self.oversized {
return Err(error(400, "INVALID_TRACE_ID"));
}
if self.len == 0 {
Ok(())
} else {
validate_trace_id(self.get())
}
}
}
impl<'de> Deserialize<'de> for Text {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct FixedText;
impl<'de> Visitor<'de> for FixedText {
type Value = Text;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a string")
}
fn visit_str<E: serde::de::Error>(self, value: &str) -> Result<Text, E> {
let mut text = Text::default();
if value.len() > MAX_ID_BYTES {
text.oversized = true;
} else {
text.bytes[..value.len()].copy_from_slice(value.as_bytes());
text.len = value.len();
}
Ok(text)
}
}
deserializer.deserialize_str(FixedText)
}
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
struct Target {
app: Text,
#[serde(rename = "interfaceId")]
interface_id: Text,
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
struct User {
#[serde(rename = "userId")]
user_id: Text,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct Trace {
#[serde(default)]
trace_id: Text,
rpc_id: Text,
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
struct Ldc {
zone: Text,
idc: Text,
env: Text,
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
struct Context {
#[serde(rename = "userInfo")]
user: User,
#[serde(rename = "traceInfo")]
trace: Trace,
#[serde(rename = "ldcInfo")]
ldc: Ldc,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct Envelope<'a> {
target: Target,
#[serde(rename = "profuseGwContext")]
context: Context,
#[serde(borrow, rename = "requestData")]
request_data: &'a RawValue,
}
#[derive(Debug)]
pub struct IngressFacts {
target: Target,
context: Context,
request_data: Range<usize>,
}
impl IngressFacts {
pub fn interface_id(&self) -> &str {
self.target.interface_id.get()
}
pub fn user_id(&self) -> &str {
self.context.user.user_id.get()
}
pub fn trace_id(&self) -> &str {
self.context.trace.trace_id.get()
}
pub fn rpc_id(&self) -> &str {
self.context.trace.rpc_id.get()
}
pub fn zone(&self) -> &str {
self.context.ldc.zone.get()
}
pub fn idc(&self) -> &str {
self.context.ldc.idc.get()
}
pub fn env(&self) -> &str {
self.context.ldc.env.get()
}
pub fn request_data_range(&self) -> Range<usize> {
self.request_data.clone()
}
}
impl ProfuseGwListenerAdapter {
pub fn recognize_observed(
&self,
method: &str,
path: &str,
content_type: &str,
body: &[u8],
parser_error: impl FnOnce(&serde_json::Error),
) -> Result<IngressFacts, CodecError> {
validate_transport(method, path, content_type, body)?;
let envelope: Envelope<'_> = serde_json::from_slice(body).map_err(|original| {
parser_error(&original);
error(400, "INVALID_JSON_ENVELOPE")
})?;
let Envelope {
target,
context,
request_data,
} = envelope;
target.app.identity()?;
target.interface_id.identity()?;
context.user.user_id.identity()?;
context.trace.trace_id.trace()?;
context.trace.rpc_id.identity()?;
context.ldc.zone.identity()?;
context.ldc.idc.identity()?;
context.ldc.env.identity()?;
let raw = request_data.get();
if !raw.starts_with('{') {
return Err(error(400, "INVALID_REQUEST_DATA"));
}
if target.app.get() != self.application.as_str() {
return Err(error(404, "APPLICATION_NOT_FOUND"));
}
let start = (raw.as_ptr() as usize)
.checked_sub(body.as_ptr() as usize)
.ok_or_else(|| error(400, "INVALID_JSON_ENVELOPE"))?;
let end = start
.checked_add(raw.len())
.filter(|end| *end <= body.len())
.ok_or_else(|| error(400, "INVALID_JSON_ENVELOPE"))?;
Ok(IngressFacts {
target,
context,
request_data: start..end,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn body(payload: &str, target: &str) -> String {
format!(
r#"{{"requestData":{payload},"profuseGwContext":{{"userInfo":{{"userId":"用户"}},"traceInfo":{{"rpcId":"0"}},"ldcInfo":{{"zone":"z","idc":"i","env":"test"}}}},"target":{target}}}"#
)
}
#[test]
fn route_at_tail_and_escaped_metadata_borrow_exact_raw_payload() {
let source = body(
r#" { "nested": [1,{"s":"a\\b"}] } "#,
r#"{"app":"app","interfaceId":"路\u7531"}"#,
);
let adapter = ProfuseGwListenerAdapter::new("app").unwrap();
let facts = adapter
.recognize_observed(METHOD, PATH, MEDIA_TYPE, source.as_bytes(), |_| {
panic!("valid JSON")
})
.unwrap();
assert_eq!(facts.interface_id(), "路由");
assert_eq!(facts.user_id(), "用户");
let raw = &source.as_bytes()[facts.request_data_range()];
assert_eq!(raw, br#"{ "nested": [1,{"s":"a\\b"}] }"#);
assert!(std::mem::size_of::<IngressFacts>() < 4096);
}
#[test]
fn bounded_recognition_preserves_codec_failures_and_original_parse_error() {
let adapter = ProfuseGwListenerAdapter::new("app").unwrap();
let target = r#"{"app":"app","interfaceId":"route"}"#;
for source in [
body("[]", target),
body("{}", r#"{"app":"other","interfaceId":"route"}"#),
body(
"{}",
&format!(r#"{{"app":"app","interfaceId":"{}"}}"#, "x".repeat(257)),
),
body("{}", r#"{"app":"app","interfaceId":"r","extra":1}"#),
body("{}", target).replace(
"\"requestData\":{}",
"\"requestData\":{},\"requestData\":{}",
),
body("{", target),
] {
let mut captured = false;
let new = adapter
.recognize_observed(METHOD, PATH, MEDIA_TYPE, source.as_bytes(), |_| {
captured = true
})
.unwrap_err();
let old = adapter
.accept(
METHOD,
PATH,
MEDIA_TYPE,
IngressIdentity::new("r", "c", 1).unwrap(),
source.as_bytes(),
)
.unwrap_err();
assert_eq!(new, old);
assert_eq!(captured, new.code == "INVALID_JSON_ENVELOPE");
}
}
}