Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
351 changes: 350 additions & 1 deletion crates/ember-core/src/schema.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,10 @@ use std::collections::HashMap;
use std::sync::{Arc, RwLock};

use bytes::Bytes;
use prost_reflect::{DescriptorPool, DynamicMessage, MessageDescriptor};
use ember_protocol::Frame;
use prost_reflect::{
DescriptorPool, DynamicMessage, FieldDescriptor, Kind, MessageDescriptor, ReflectMessage,
};
use thiserror::Error;

/// Errors that can occur during schema operations.
Expand All @@ -28,6 +31,9 @@ pub enum SchemaError {

#[error("schema already registered: {0}")]
AlreadyExists(String),

#[error("field not found: {0}")]
FieldNotFound(String),
}

/// A registered schema: the raw descriptor bytes and the parsed pool.
Expand Down Expand Up @@ -169,6 +175,25 @@ impl SchemaRegistry {
);
}

/// Reads a single field from an encoded protobuf message.
///
/// Decodes the message using the schema registry, walks the dot-separated
/// `field_path` to the target field, and converts the value to a RESP3
/// frame. Returns an error for complex types (message, list, map) — those
/// require `PROTO.GET` for full deserialization.
pub fn get_field(
&self,
type_name: &str,
data: &[u8],
field_path: &str,
) -> Result<Frame, SchemaError> {
let descriptor = self.find_message(type_name)?;
let msg = DynamicMessage::decode(descriptor, data)
.map_err(|e| SchemaError::ValidationFailed(e.to_string()))?;
let (value, field_desc) = resolve_field_path(&msg, field_path)?;
value_to_frame(&value, &field_desc)
}

/// Looks up a message descriptor by full name across all schemas.
fn find_message(&self, message_type: &str) -> Result<MessageDescriptor, SchemaError> {
for schema in self.schemas.values() {
Expand All @@ -180,6 +205,105 @@ impl SchemaRegistry {
}
}

/// Walks a dot-separated field path through a `DynamicMessage`, returning
/// the leaf value (owned) and its field descriptor.
///
/// Intermediate path segments must be message-typed fields. The leaf
/// segment is the target field whose value is returned.
fn resolve_field_path(
msg: &DynamicMessage,
path: &str,
) -> Result<(prost_reflect::Value, FieldDescriptor), SchemaError> {
if path.is_empty() {
return Err(SchemaError::FieldNotFound("empty field path".into()));
}

let segments: Vec<&str> = path.split('.').collect();
for seg in &segments {
if seg.is_empty() {
return Err(SchemaError::FieldNotFound(format!(
"invalid field path '{path}': empty segment"
)));
}
}

let mut current_msg = msg.clone();

for (i, segment) in segments.iter().enumerate() {
let field_desc = current_msg
.descriptor()
.get_field_by_name(segment)
.ok_or_else(|| SchemaError::FieldNotFound(segment.to_string()))?;

let value = current_msg.get_field(&field_desc).into_owned();

if i == segments.len() - 1 {
return Ok((value, field_desc));
}

// intermediate segment — must be a message type
match value {
prost_reflect::Value::Message(nested) => {
current_msg = nested;
}
_ => {
return Err(SchemaError::FieldNotFound(format!(
"'{segment}' is not a message field, cannot traverse further"
)));
}
}
}

unreachable!("loop always returns at the leaf segment")
}

/// Converts a `prost_reflect::Value` + its field descriptor into a RESP3 frame.
///
/// Scalar types are mapped to native RESP3 types. Complex types (message,
/// repeated, map) return an error directing clients to use `PROTO.GET`.
fn value_to_frame(
value: &prost_reflect::Value,
field_desc: &FieldDescriptor,
) -> Result<Frame, SchemaError> {
// reject repeated and map fields up front
if field_desc.is_list() || field_desc.is_map() {
return Err(SchemaError::ValidationFailed(
"use PROTO.GET for repeated/map fields".into(),
));
}

match value {
prost_reflect::Value::String(s) => Ok(Frame::Bulk(Bytes::from(s.clone()))),
prost_reflect::Value::Bytes(b) => Ok(Frame::Bulk(b.clone())),
prost_reflect::Value::I32(n) => Ok(Frame::Integer(i64::from(*n))),
prost_reflect::Value::I64(n) => Ok(Frame::Integer(*n)),
prost_reflect::Value::U32(n) => Ok(Frame::Integer(i64::from(*n))),
prost_reflect::Value::U64(n) => Ok(Frame::Integer(*n as i64)),
prost_reflect::Value::F32(n) => Ok(Frame::Bulk(Bytes::from(format!("{n}")))),
prost_reflect::Value::F64(n) => Ok(Frame::Bulk(Bytes::from(format!("{n}")))),
prost_reflect::Value::Bool(b) => Ok(Frame::Integer(if *b { 1 } else { 0 })),
prost_reflect::Value::EnumNumber(n) => {
// look up the enum value name from the descriptor
if let Kind::Enum(enum_desc) = field_desc.kind() {
if let Some(val) = enum_desc.get_value(*n) {
return Ok(Frame::Bulk(Bytes::from(val.name().to_owned())));
}
}
// fallback: return the numeric value
Ok(Frame::Integer(i64::from(*n)))
}
prost_reflect::Value::Message(_) => Err(SchemaError::ValidationFailed(
"use PROTO.GET for nested message fields".into(),
)),
prost_reflect::Value::List(_) => Err(SchemaError::ValidationFailed(
"use PROTO.GET for repeated fields".into(),
)),
prost_reflect::Value::Map(_) => Err(SchemaError::ValidationFailed(
"use PROTO.GET for map fields".into(),
)),
}
}

impl std::fmt::Debug for SchemaRegistry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SchemaRegistry")
Expand Down Expand Up @@ -334,4 +458,229 @@ mod tests {
pairs.sort();
assert_eq!(pairs, vec!["alpha", "beta"]);
}

// --- get_field tests ---

/// Helper: encode a test.User message with the given name.
fn encode_user(registry: &SchemaRegistry, name: &str) -> Vec<u8> {
let pool = &registry.schemas["users"].pool;
let msg_desc = pool.get_message_by_name("test.User").unwrap();
let mut msg = DynamicMessage::new(msg_desc);
msg.set_field_by_name("name", prost_reflect::Value::String(name.into()));
let mut buf = Vec::new();
use prost_reflect::prost::Message;
msg.encode(&mut buf).unwrap();
buf
}

#[test]
fn get_field_string() {
let mut registry = SchemaRegistry::new();
let desc = make_descriptor("test", "User", "name");
registry.register("users".into(), desc).unwrap();

let data = encode_user(&registry, "alice");
let frame = registry.get_field("test.User", &data, "name").unwrap();
assert_eq!(frame, Frame::Bulk(Bytes::from("alice")));
}

#[test]
fn get_field_default_value() {
let mut registry = SchemaRegistry::new();
let desc = make_descriptor("test", "User", "name");
registry.register("users".into(), desc).unwrap();

// encode an empty message (no fields set)
let pool = &registry.schemas["users"].pool;
let msg_desc = pool.get_message_by_name("test.User").unwrap();
let msg = DynamicMessage::new(msg_desc);
let mut buf = Vec::new();
use prost_reflect::prost::Message;
msg.encode(&mut buf).unwrap();

// default string should be empty
let frame = registry.get_field("test.User", &buf, "name").unwrap();
assert_eq!(frame, Frame::Bulk(Bytes::from("")));
}

#[test]
fn get_field_int() {
use prost_reflect::prost_types::{
DescriptorProto, FieldDescriptorProto, FileDescriptorProto, FileDescriptorSet,
};

let fds = FileDescriptorSet {
file: vec![FileDescriptorProto {
name: Some("test.proto".into()),
package: Some("test".into()),
message_type: vec![DescriptorProto {
name: Some("Counter".into()),
field: vec![FieldDescriptorProto {
name: Some("count".into()),
number: Some(1),
r#type: Some(5), // TYPE_INT32
label: Some(1),
..Default::default()
}],
..Default::default()
}],
..Default::default()
}],
};
let mut desc_buf = Vec::new();
use prost_reflect::prost::Message;
fds.encode(&mut desc_buf).unwrap();
let desc = Bytes::from(desc_buf);

let mut registry = SchemaRegistry::new();
registry.register("counters".into(), desc.clone()).unwrap();

let pool = &registry.schemas["counters"].pool;
let msg_desc = pool.get_message_by_name("test.Counter").unwrap();
let mut msg = DynamicMessage::new(msg_desc);
msg.set_field_by_name("count", prost_reflect::Value::I32(42));
let mut buf = Vec::new();
msg.encode(&mut buf).unwrap();

let frame = registry.get_field("test.Counter", &buf, "count").unwrap();
assert_eq!(frame, Frame::Integer(42));
}

#[test]
fn get_field_bool() {
use prost_reflect::prost_types::{
DescriptorProto, FieldDescriptorProto, FileDescriptorProto, FileDescriptorSet,
};

let fds = FileDescriptorSet {
file: vec![FileDescriptorProto {
name: Some("test.proto".into()),
package: Some("test".into()),
message_type: vec![DescriptorProto {
name: Some("Flag".into()),
field: vec![FieldDescriptorProto {
name: Some("active".into()),
number: Some(1),
r#type: Some(8), // TYPE_BOOL
label: Some(1),
..Default::default()
}],
..Default::default()
}],
..Default::default()
}],
};
let mut desc_buf = Vec::new();
use prost_reflect::prost::Message;
fds.encode(&mut desc_buf).unwrap();
let desc = Bytes::from(desc_buf);

let mut registry = SchemaRegistry::new();
registry.register("flags".into(), desc).unwrap();

let pool = &registry.schemas["flags"].pool;
let msg_desc = pool.get_message_by_name("test.Flag").unwrap();
let mut msg = DynamicMessage::new(msg_desc);
msg.set_field_by_name("active", prost_reflect::Value::Bool(true));
let mut buf = Vec::new();
msg.encode(&mut buf).unwrap();

let frame = registry.get_field("test.Flag", &buf, "active").unwrap();
assert_eq!(frame, Frame::Integer(1));
}

/// Builds a descriptor with a nested message: Outer { Inner inner = 1; }
/// where Inner { string value = 1; }
fn make_nested_descriptor() -> Bytes {
use prost_reflect::prost_types::{
DescriptorProto, FieldDescriptorProto, FileDescriptorProto, FileDescriptorSet,
};

let fds = FileDescriptorSet {
file: vec![FileDescriptorProto {
name: Some("test.proto".into()),
package: Some("test".into()),
message_type: vec![
DescriptorProto {
name: Some("Inner".into()),
field: vec![FieldDescriptorProto {
name: Some("value".into()),
number: Some(1),
r#type: Some(9), // TYPE_STRING
label: Some(1),
..Default::default()
}],
..Default::default()
},
DescriptorProto {
name: Some("Outer".into()),
field: vec![FieldDescriptorProto {
name: Some("inner".into()),
number: Some(1),
r#type: Some(11), // TYPE_MESSAGE
label: Some(1),
type_name: Some(".test.Inner".into()),
..Default::default()
}],
..Default::default()
},
],
..Default::default()
}],
};
let mut buf = Vec::new();
use prost_reflect::prost::Message;
fds.encode(&mut buf).unwrap();
Bytes::from(buf)
}

#[test]
fn get_field_nested_path() {
let desc = make_nested_descriptor();
let mut registry = SchemaRegistry::new();
registry.register("nested".into(), desc).unwrap();

let pool = &registry.schemas["nested"].pool;
let outer_desc = pool.get_message_by_name("test.Outer").unwrap();
let inner_desc = pool.get_message_by_name("test.Inner").unwrap();

let mut inner = DynamicMessage::new(inner_desc);
inner.set_field_by_name("value", prost_reflect::Value::String("hello".into()));

let mut outer = DynamicMessage::new(outer_desc);
outer.set_field_by_name("inner", prost_reflect::Value::Message(inner));

let mut buf = Vec::new();
use prost_reflect::prost::Message;
outer.encode(&mut buf).unwrap();

let frame = registry
.get_field("test.Outer", &buf, "inner.value")
.unwrap();
assert_eq!(frame, Frame::Bulk(Bytes::from("hello")));
}

#[test]
fn get_field_nonexistent() {
let mut registry = SchemaRegistry::new();
let desc = make_descriptor("test", "User", "name");
registry.register("users".into(), desc).unwrap();

let data = encode_user(&registry, "alice");
let err = registry
.get_field("test.User", &data, "nonexistent")
.unwrap_err();
assert!(matches!(err, SchemaError::FieldNotFound(_)));
}

#[test]
fn get_field_empty_path() {
let mut registry = SchemaRegistry::new();
let desc = make_descriptor("test", "User", "name");
registry.register("users".into(), desc).unwrap();

let data = encode_user(&registry, "alice");
let err = registry.get_field("test.User", &data, "").unwrap_err();
assert!(matches!(err, SchemaError::FieldNotFound(_)));
}
}
Loading
Loading