From c2a7afe6c19fb1e6c3c0806c67f4b3c189b6bc4c Mon Sep 17 00:00:00 2001 From: Alex Emmet <111742636+Alex-Emmet@users.noreply.github.com> Date: Mon, 22 Jun 2026 23:29:18 +0200 Subject: [PATCH] Values, Cleaning, Docs, Tests, Example (Current Example is Wrong) --- CONNECTOR.md | 10 +- Cargo.toml | 46 +++++- README.md | 50 ++++-- TYPE-MAP.md | 15 +- client/Cargo.toml | 5 + client/src/lib.rs | 256 ++++++++++++++++++++++++++++++- codec/Cargo.toml | 9 +- codec/src/communication_value.rs | 134 +++++++++++----- codec/src/data_value.rs | 25 +-- codec/src/lib.rs | 16 +- codec/src/registry.rs | 21 +-- codec/src/util.rs | 9 -- common/Cargo.toml | 4 +- common/src/lib.rs | 29 ++-- crypto/Cargo.toml | 2 +- crypto/src/helper.rs | 47 +----- crypto/src/kem.rs | 2 +- crypto/src/keypair.rs | 36 ++--- crypto/src/lib.rs | 246 ++++++++++++++++++++++++++++- crypto/src/sign.rs | 102 ++++++------ example-type-maps.yaml | 8 +- example-usage/.gitignore | 4 + example-usage/client/Cargo.toml | 5 +- example-usage/client/src/main.rs | 109 ++++++++++--- example-usage/server/Cargo.toml | 5 +- example-usage/server/src/main.rs | 203 ++++++++++++++++++------ host/Cargo.toml | 5 + host/src/lib.rs | 253 +++++++++++++++++++++++++++++- registry/Cargo.toml | 8 - registry/src/lib.rs | 127 --------------- src/lib.rs | 13 ++ transport/Cargo.toml | 5 +- transport/src/connection.rs | 10 +- transport/src/host.rs | 3 +- type-map/Cargo.toml | 4 +- type-map/build.rs | 244 +++++++++++++++++++++++------ type-map/src/lib.rs | 137 +++++++++++++++++ 37 files changed, 1690 insertions(+), 517 deletions(-) delete mode 100644 codec/src/util.rs create mode 100644 example-usage/.gitignore delete mode 100644 registry/Cargo.toml delete mode 100644 registry/src/lib.rs create mode 100644 src/lib.rs diff --git a/CONNECTOR.md b/CONNECTOR.md index be680d0..f91196f 100644 --- a/CONNECTOR.md +++ b/CONNECTOR.md @@ -4,10 +4,10 @@ This file documents the Connection and Version Negotiation logic. ## Registry -The `registry` crate provides a multi-version `Registry` used by the host for version negotiation. A `Registry` holds one `TypeMap` per protocol version and supports `negotiate()`: +The `registry` crate provides a multi-version `Registry` used by the host for version negotiation. Accessed through the `mtp` facade (requires the `host` feature): ```rust -use registry::Registry; +use mtp::codec::registry::Registry; let registry = Registry::builtin(); // loads all TypeMaps from config @@ -34,7 +34,7 @@ The host creates a QUIC server, manages the registry, and handles version negoti ### Initialization ```rust -use mtp_host::{MTPHost, HostConfig}; +use mtp::host::{MTPHost, HostConfig}; let config = HostConfig { ip: "::".into(), @@ -87,7 +87,7 @@ New clients use the `Register` variant instead, presenting their public key for The client connects to a host and uses a single compiled-in protocol version. ```rust -use mtp_client::{MTPClient, ClientConfig}; +use mtp::client::{MTPClient, ClientConfig}; let config = ClientConfig { url: "https://host.example.com:4433".into(), @@ -103,7 +103,7 @@ let conn = MTPClient::connect(config, 8765).await?; */ ``` -The client's `PROTOCOL_VERSION` constant is set by `protocol_version` in `type-maps.yaml` and baked in at compile time. The client never imports the `registry` crate; it only uses `type-map` for enum types and `codec` for encoding. +The client's `PROTOCOL_VERSION` constant is set by `protocol_version` in `type-maps.yaml` and baked in at compile time. The client never imports the `registry` crate; it only uses `mtp::type_map` for enum types and `mtp::codec` for encoding. --- diff --git a/Cargo.toml b/Cargo.toml index 78c22f3..2627783 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -2,7 +2,6 @@ members = [ "common", "crypto", - "registry", "type-map", "codec", "transport", @@ -10,3 +9,48 @@ members = [ "client", ] resolver = "3" + +# ============================================================================= +# mtp - umbrella crate +# +# Re-exports all MTP sub-crates behind feature flags so consumers write: +# +# [dependencies] +# mtp = { version = "0.1", features = ["client", "crypto"] } +# +# instead of listing every sub-crate individually. +# ============================================================================= +[package] +name = "mtp" +version = "0.1.0" +edition = "2024" + +[dependencies] +# --- always-on core --- +mtp-common = { path = "common" } +mtp-type-map = { path = "type-map" } +mtp-codec = { path = "codec" } +mtp-transport = { path = "transport" } + +# --- optional, behind features --- +mtp-crypto = { path = "crypto", optional = true } +mtp-host = { path = "host", optional = true } +mtp-client = { path = "client", optional = true } + +[features] +default = [] + +# Message encryption (AEAD, signatures, KEM, KDF, hashing). +# When combined with `host` or `client`, also enables connection authentication. +crypto = [ + "dep:mtp-crypto", + "mtp-codec/crypto", + "mtp-host?/crypto", + "mtp-client?/crypto", +] + +# MTP server host - version negotiation, Registry, incoming QUIC connections. +host = ["dep:mtp-host", "mtp-codec/registry", "mtp-transport/host"] + +# MTP client - outgoing QUIC connections to a host. +client = ["dep:mtp-client"] diff --git a/README.md b/README.md index 582ac01..57e2412 100644 --- a/README.md +++ b/README.md @@ -2,17 +2,41 @@ **MTP** is a **m**odular **t**ransport **p**rotocol by Methanium. -## Crates -- Codec -- Transport -- Common -- Crypto -- TypeMap -- Registry -- Host -- Client +## Getting Started -## Codec +Add the `mtp` crate with your desired features: + +```toml +[dependencies] +mtp = { path = "..", features = ["client", "crypto"] } +``` + +The `mtp` umbrella crate re-exports all sub-crates behind feature flags: + +| Feature | Pulls in | Enables | +| -------- | --------------------------- | ----------------------------------------- | +| `crypto` | `mtp::crypto` | AEAD, signatures, KEM, KDF, hashing | +| `host` | `mtp::host`, `mtp::codec::registry` | QUIC server, version negotiation | +| `client` | `mtp::client` | QUIC client connections | + +Core crates (`codec`, `transport`, `common`, `type_map`) are always available. + +```rust +use mtp::codec::{encode, decode, DataValue}; +use mtp::type_map::TypeMap; +use mtp::transport::Sender; + +#[cfg(feature = "crypto")] +use mtp::crypto::ChaCha20Poly1305; +``` + +## Sub-crates + +All sub-crates are re-exported through the `mtp` facade and can be referenced +as `mtp::codec`, `mtp::transport`, `mtp::common`, `mtp::type_map`, +`mtp::crypto`, `mtp::host`, `mtp::client`. + +### Codec The Codec crate handles encoding and decoding of MTP packets using Communication Types and Data Types from the Registry. @@ -33,15 +57,15 @@ The Codec uses the Crypto crate to encrypt and decrypt Encrypted Containers. --- -## Transport +### Transport The Transport crate wraps QUIC using `wtransport`. It provides `Sender`/`Receiver` for bidirectional message passing over QUIC streams. -## Common +### Common Common handles logging and error handling. Personal information reported to Common will be anonymized. -## Crypto Stack +### Crypto Stack | Crate | Audited? | Notes | | ---------------- | -------- | ------------------------------------------------ | diff --git a/TYPE-MAP.md b/TYPE-MAP.md index 41342f2..ce650f3 100644 --- a/TYPE-MAP.md +++ b/TYPE-MAP.md @@ -17,7 +17,7 @@ An example `type-maps.yaml` is provided in the [`example-type-maps.yaml`](./exam After editing the config and rebuilding, `CommunicationType` and `DataType` enums are generated automatically. Use them in code: ```rust -use mtp_type_map::{CommunicationType, DataType, TypeMap}; +use mtp::type_map::{CommunicationType, DataType, TypeMap}; let tm = TypeMap::v2_0(); let id = tm.data_id_enum(DataType::SomeType).unwrap(); @@ -28,8 +28,8 @@ The enums are a **union across all versions**; every type name from every versio Encoding/decoding uses a `TypeMap` to resolve type names to wire IDs: ```rust -use mtp_codec::{encode, decode, DataValue}; -use mtp_type_map::TypeMap; +use mtp::codec::{encode, decode, DataValue}; +use mtp::type_map::TypeMap; let tm = TypeMap::v2_0(); let value = DataValue::Str("hello".into()); @@ -62,8 +62,15 @@ This is by design: the host maps unknown types to `Error`, and the client should The `registry` feature of the Codec crate adds `VersionedCodec` for version-aware encoding: +Requires the `host` feature (which enables `mtp-codec`'s `registry` feature): + +```toml +[dependencies] +mtp = { path = "..", features = ["host"] } +``` + ```rust -use mtp_codec::registry::{Registry, VersionedCodec}; +use mtp::codec::registry::{Registry, VersionedCodec}; let registry = Registry::builtin(); let codec = VersionedCodec::new(registry); diff --git a/client/Cargo.toml b/client/Cargo.toml index ee1bbb5..770adb8 100644 --- a/client/Cargo.toml +++ b/client/Cargo.toml @@ -7,4 +7,9 @@ edition = "2024" mtp-common = { path = "../common" } mtp-codec = { path = "../codec" } mtp-transport = { path = "../transport" } +mtp-crypto = { path = "../crypto", optional = true } tokio = { version = "1", features = ["full"] } +rand = "0.8" + +[features] +crypto = ["dep:mtp-crypto", "mtp-codec/crypto"] diff --git a/client/src/lib.rs b/client/src/lib.rs index 6c87d46..4d030d2 100644 --- a/client/src/lib.rs +++ b/client/src/lib.rs @@ -1,4 +1,6 @@ -use mtp_codec::{CommunicationValue, DataTypeId, DataValue, Version, PROTOCOL_VERSION}; +#[cfg(feature = "crypto")] +use mtp_codec::DataType; +use mtp_codec::{CommunicationValue, DataTypeId, DataValue, PROTOCOL_VERSION, Version}; use mtp_common::CommunicationError; use mtp_transport::{Policy, Receiver, Sender}; @@ -12,6 +14,18 @@ pub struct MTPConnection { pub version: Version, pub sender: Sender, pub receiver: Receiver, + #[cfg(feature = "crypto")] + pub auth_state: AuthState, + #[cfg(feature = "crypto")] + pub client_id: u64, +} + +#[cfg(feature = "crypto")] +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum AuthState { + Pending, + Authenticated, + Failed, } pub struct MTPClient; @@ -42,6 +56,246 @@ impl MTPClient { version: PROTOCOL_VERSION, sender, receiver, + #[cfg(feature = "crypto")] + auth_state: AuthState::Authenticated, + #[cfg(feature = "crypto")] + client_id, }) } } + +/* ===== Authentication ===== */ +#[cfg(feature = "crypto")] +impl MTPClient { + pub async fn auth_connect( + config: ClientConfig, + client_id: u64, + keys: mtp_crypto::KeyGroup, + host_public_key: mtp_crypto::SignaturePublicKey, + ) -> Result { + use mtp_crypto::{Ed25519Signer, SignatureScheme, verify_ed25519}; + + let (sender, receiver) = + mtp_transport::connect(&config.url, config.server_cert, Policy::default()).await?; + + // 1. Receive host greeting (IdentificationResponse with host ID, public keys, nonce) + let _greeting = receiver.receive().await?; + + // 2. Build and send Identification message with auth data + let client_nonce: u128 = rand::random(); + let version_str = format!("{}", PROTOCOL_VERSION); + + let mut sig_payload = Vec::new(); + sig_payload.extend_from_slice(version_str.as_bytes()); + sig_payload.extend_from_slice(&client_id.to_be_bytes()); + sig_payload.extend_from_slice(&client_nonce.to_be_bytes()); + + let signer = Ed25519Signer::new(&keys.signature_private_key) + .map_err(|e| CommunicationError::Other(e.to_string()))?; + + /* ===== Signature ===== */ + // Current: signing concatenated bytes + // Future: sign entire CommunicationValue + let signature = signer + .sign(&sig_payload) + .map_err(|e| CommunicationError::Other(e.to_string()))?; + /* ===== End Signature ===== */ + + let ident = CommunicationValue::new(mtp_codec::CommunicationType::Identification) + .add_typed_default(DataType::Version, DataValue::Str(version_str)) + .add_typed_default(DataType::Id, DataValue::UnsignedNumber(client_id as u128)) + .add_typed_default( + DataType::ClientNonce, + DataValue::UnsignedNumber(client_nonce), + ) + .add_typed_default(DataType::Signature, DataValue::Bytes(signature)); + + sender.send(&ident).await?; + + // 3. Receive host response + let response = receiver.receive().await?; + + let connected = response.get_data(DataTypeId(11)); + match connected { + DataValue::BoolTrue => {} + DataValue::BoolFalse => { + return Err(CommunicationError::AuthenticationFailed( + "Server rejected authentication".into(), + )); + } + _ => { + return Err(CommunicationError::AuthenticationFailed( + "Invalid response".into(), + )); + } + } + + let echo_nonce = response.get_data(DataTypeId(7)); + match echo_nonce { + DataValue::UnsignedNumber(n) if *n == client_nonce as u128 => {} + _ => { + return Err(CommunicationError::AuthenticationFailed( + "Nonce mismatch".into(), + )); + } + } + + let host_new_nonce = match response.get_data(DataTypeId(5)) { + DataValue::UnsignedNumber(n) => *n as u128, + _ => { + return Err(CommunicationError::AuthenticationFailed( + "Missing new nonce".into(), + )); + } + }; + + let host_sig = match response.get_data(DataTypeId(10)) { + DataValue::Bytes(b) => b.clone(), + _ => { + return Err(CommunicationError::AuthenticationFailed( + "Missing signature".into(), + )); + } + }; + + let mut host_sig_payload = Vec::new(); + host_sig_payload.extend_from_slice(b"true"); + host_sig_payload.extend_from_slice(&client_nonce.to_be_bytes()); + host_sig_payload.extend_from_slice(&host_new_nonce.to_be_bytes()); + + /* ===== Signature ===== */ + verify_ed25519(&host_public_key, &host_sig_payload, &host_sig).map_err(|_| { + CommunicationError::AuthenticationFailed("Host signature invalid".into()) + })?; + /* ===== End Signature ===== */ + + Ok(MTPConnection { + version: PROTOCOL_VERSION, + sender, + receiver, + auth_state: AuthState::Authenticated, + client_id, + }) + } + + pub async fn auth_register( + config: ClientConfig, + keys: mtp_crypto::Keyring, + host_public_key: mtp_crypto::SignaturePublicKey, + ) -> Result { + use mtp_crypto::{Ed25519Signer, SignatureScheme, verify_ed25519}; + + let (sender, receiver) = + mtp_transport::connect(&config.url, config.server_cert, Policy::default()).await?; + + // 1. Receive host greeting (IdentificationResponse with host ID, public keys, nonce) + let _greeting = receiver.receive().await?; + + // 2. Build and send Register message with public keys and auth data + let client_nonce: u128 = rand::random(); + let version_str = format!("{}", PROTOCOL_VERSION); + let pk_bundle = keys.public_key_bundle(); + let pk_bytes = pk_bundle.as_bytes(); + + let mut sig_payload = Vec::new(); + sig_payload.extend_from_slice(version_str.as_bytes()); + sig_payload.extend_from_slice(&client_nonce.to_be_bytes()); + sig_payload.extend_from_slice(&pk_bytes); + + let signer = Ed25519Signer::new(&keys.sig_cl_secret_key) + .map_err(|e| CommunicationError::Other(e.to_string()))?; + + /* ===== Signature ===== */ + let signature = signer + .sign(&sig_payload) + .map_err(|e| CommunicationError::Other(e.to_string()))?; + /* ===== End Signature ===== */ + + let register = CommunicationValue::new(mtp_codec::CommunicationType::Register) + .add_typed_default(DataType::Version, DataValue::Str(version_str)) + .add_typed_default( + DataType::ClientNonce, + DataValue::UnsignedNumber(client_nonce), + ) + .add_typed_default(DataType::PublicKeys, DataValue::Bytes(pk_bytes)) + .add_typed_default(DataType::Signature, DataValue::Bytes(signature)); + + sender.send(®ister).await?; + + // 3. Receive host response + let response = receiver.receive().await?; + + let connected = response.get_data(DataTypeId(11)); + match connected { + DataValue::BoolTrue => {} + DataValue::BoolFalse => { + return Err(CommunicationError::AuthenticationFailed( + "Server rejected registration".into(), + )); + } + _ => { + return Err(CommunicationError::AuthenticationFailed( + "Invalid response".into(), + )); + } + } + + let assigned_id = match response.get_data(DataTypeId(6)) { + DataValue::UnsignedNumber(n) => *n as u128, + _ => { + return Err(CommunicationError::AuthenticationFailed( + "Missing assigned ID".into(), + )); + } + }; + + let echo_nonce = response.get_data(DataTypeId(7)); + match echo_nonce { + DataValue::UnsignedNumber(n) if *n == client_nonce as u128 => {} + _ => { + return Err(CommunicationError::AuthenticationFailed( + "Nonce mismatch".into(), + )); + } + } + + let host_new_nonce = match response.get_data(DataTypeId(5)) { + DataValue::UnsignedNumber(n) => *n as u128, + _ => { + return Err(CommunicationError::AuthenticationFailed( + "Missing new nonce".into(), + )); + } + }; + + let host_sig = match response.get_data(DataTypeId(10)) { + DataValue::Bytes(b) => b.clone(), + _ => { + return Err(CommunicationError::AuthenticationFailed( + "Missing signature".into(), + )); + } + }; + + let mut host_sig_payload = Vec::new(); + host_sig_payload.extend_from_slice(b"true"); + host_sig_payload.extend_from_slice(&assigned_id.to_be_bytes()); + host_sig_payload.extend_from_slice(&client_nonce.to_be_bytes()); + host_sig_payload.extend_from_slice(&host_new_nonce.to_be_bytes()); + + /* ===== Signature ===== */ + verify_ed25519(&host_public_key, &host_sig_payload, &host_sig).map_err(|_| { + CommunicationError::AuthenticationFailed("Host signature invalid".into()) + })?; + /* ===== End Signature ===== */ + + Ok(MTPConnection { + version: PROTOCOL_VERSION, + sender, + receiver, + auth_state: AuthState::Authenticated, + client_id: assigned_id as u64, + }) + } +} +/* ===== End Authentication ===== */ diff --git a/codec/Cargo.toml b/codec/Cargo.toml index 4b75c72..6e91169 100644 --- a/codec/Cargo.toml +++ b/codec/Cargo.toml @@ -7,12 +7,11 @@ edition = "2024" mtp-type-map = { path = "../type-map" } mtp-common = { path = "../common" } mtp-crypto = { path = "../crypto", optional = true } -registry = { path = "../registry", optional = true } -base64 = "*" -byteorder = "*" -rand = { version = "*", features = ["std", "std_rng"] } +base64 = "0.22" +byteorder = "1.5" +rand = { version = "0.8", features = ["std", "std_rng"] } [features] default = [] -registry = ["dep:registry", "mtp-type-map/multi-version"] +registry = ["mtp-type-map/registry"] crypto = ["dep:mtp-crypto"] diff --git a/codec/src/communication_value.rs b/codec/src/communication_value.rs index 68bcaf3..db3733c 100644 --- a/codec/src/communication_value.rs +++ b/codec/src/communication_value.rs @@ -5,6 +5,7 @@ use std::io::{Cursor, Read}; use crate::data_value::DataValue; use crate::rand_u32; +use mtp_common::CodecError; use mtp_type_map::{ communication_type_name, data_type_name, CommunicationType, CommunicationTypeId, DataType, DataTypeId, TypeMap, PROTOCOL_VERSION, @@ -20,8 +21,8 @@ pub struct CommunicationValue { type_map: Option, } -#[allow(dead_code)] impl CommunicationValue { + #[must_use] pub fn new(comm_type: CommunicationType) -> Self { let tm = TypeMap::new(PROTOCOL_VERSION); let id = comm_type.to_id(&tm); @@ -36,6 +37,7 @@ impl CommunicationValue { } #[cfg(feature = "registry")] + #[must_use] pub fn from_comm(comm_type: CommunicationType, tm: &TypeMap) -> Self { let id = comm_type.to_id(tm); Self { @@ -48,19 +50,7 @@ impl CommunicationValue { } } - pub fn from_comm_default(comm_type: CommunicationType) -> Self { - let tm = TypeMap::new(PROTOCOL_VERSION); - let id = comm_type.to_id(&tm); - Self { - id: rand_u32(), - comm_type: id, - sender: 0, - receiver: 0, - data: BTreeMap::new(), - type_map: Some(tm), - } - } - + #[must_use] pub fn with_id(mut self, p0: u32) -> Self { self.id = p0; self @@ -70,6 +60,7 @@ impl CommunicationValue { self.id } + #[must_use] pub fn with_sender(mut self, sender: u64) -> Self { self.sender = sender; self @@ -79,6 +70,7 @@ impl CommunicationValue { self.sender } + #[must_use] pub fn with_receiver(mut self, receiver: u64) -> Self { self.receiver = receiver; self @@ -92,21 +84,20 @@ impl CommunicationValue { self.comm_type } - pub fn is_type(&self, p0: CommunicationTypeId) -> bool { - self.comm_type == p0 - } - + #[must_use] pub fn add_data(mut self, data: DataTypeId, value: DataValue) -> Self { self.data.insert(data, value); self } #[cfg(feature = "registry")] + #[must_use] pub fn add_typed(mut self, data: DataType, tm: &TypeMap, value: DataValue) -> Self { self.data.insert(data.to_id(tm), value); self } + #[must_use] pub fn add_typed_default(mut self, data: DataType, value: DataValue) -> Self { let tm = self.type_map.clone().unwrap_or_else(TypeMap::latest); self.data.insert(data.to_id(&tm), value); @@ -117,9 +108,6 @@ impl CommunicationValue { self.data.get(&data_type).unwrap_or(&DataValue::Null) } - pub fn get_data_container(&self) -> &BTreeMap { - &self.data - } } impl CommunicationValue { @@ -184,33 +172,43 @@ impl CommunicationValue { frame } - pub fn from_bytes(bytes: &[u8]) -> Option { + pub fn from_bytes(bytes: &[u8]) -> Result { let mut cursor = Cursor::new(bytes); - let total_len = cursor.read_u32::().ok()? as usize; + let total_len = cursor + .read_u32::() + .map_err(|_| CodecError::InvalidEncoding)? as usize; if bytes.len() < 4 + total_len { - return None; + return Err(CodecError::InvalidEncoding); } let frame_end = 4 + total_len; - let comm_type_num = cursor.read_u16::().ok()?; + let comm_type_num = cursor + .read_u16::() + .map_err(|_| CodecError::InvalidEncoding)?; let comm_type = CommunicationTypeId(comm_type_num); - let flags = cursor.read_u8().ok()?; + let flags = cursor + .read_u8() + .map_err(|_| CodecError::InvalidEncoding)?; let has_sender = (flags & 0b0000_0001) != 0; let has_receiver = (flags & 0b0000_0010) != 0; let has_id = (flags & 0b0000_0100) != 0; let id = if has_id { - cursor.read_u32::().ok()? + cursor + .read_u32::() + .map_err(|_| CodecError::InvalidEncoding)? } else { 0 }; let sender = if has_sender { let mut buf = [0u8; 8]; - cursor.read_exact(&mut buf[2..]).ok()?; + cursor + .read_exact(&mut buf[2..]) + .map_err(|_| CodecError::InvalidEncoding)?; u64::from_be_bytes(buf) } else { 0 @@ -218,7 +216,9 @@ impl CommunicationValue { let receiver = if has_receiver { let mut buf = [0u8; 8]; - cursor.read_exact(&mut buf[2..]).ok()?; + cursor + .read_exact(&mut buf[2..]) + .map_err(|_| CodecError::InvalidEncoding)?; u64::from_be_bytes(buf) } else { 0 @@ -226,14 +226,14 @@ impl CommunicationValue { let pos = cursor.position() as usize; if pos > frame_end { - return None; + return Err(CodecError::InvalidEncoding); } let data_bytes = &bytes[pos..frame_end]; - let data_value = DataValue::from_bytes(data_bytes)?; - let data = data_value.as_map()?; + let data_value = DataValue::from_bytes(data_bytes).ok_or(CodecError::InvalidEncoding)?; + let data = data_value.as_map().ok_or(CodecError::InvalidEncoding)?; - Some(Self { + Ok(Self { id, comm_type, sender, @@ -242,6 +242,51 @@ impl CommunicationValue { type_map: None, }) } + + pub fn from_bytes_with(bytes: &[u8], tm: &TypeMap) -> Result { + let mut val = Self::from_bytes(bytes)?; + val.type_map = Some(tm.clone()); + Ok(val) + } + + #[cfg(feature = "registry")] + pub fn migrate(&self, target_tm: &TypeMap) -> Result { + let comm_name = + communication_type_name(self.comm_type.0).ok_or_else(|| { + CodecError::UnknownCommunicationType(self.comm_type.0.to_string()) + })?; + let comm_variant = CommunicationType::from_name(comm_name).ok_or_else(|| { + CodecError::UnknownCommunicationType(comm_name.to_string()) + })?; + let new_comm_id = CommunicationTypeId( + target_tm + .comm_id_enum(comm_variant) + .ok_or_else(|| CodecError::UnknownCommunicationType(comm_name.to_string()))?, + ); + + let mut new_data = BTreeMap::new(); + for (&old_id, value) in &self.data { + let name = data_type_name(old_id.0) + .ok_or_else(|| CodecError::UnknownDataType(old_id.0.to_string()))?; + let variant = DataType::from_name(name) + .ok_or_else(|| CodecError::UnknownDataType(name.to_string()))?; + let new_id = DataTypeId( + target_tm + .data_id_enum(variant) + .ok_or_else(|| CodecError::UnknownDataType(name.to_string()))?, + ); + new_data.insert(new_id, value.clone()); + } + + Ok(Self { + id: self.id, + comm_type: new_comm_id, + sender: self.sender, + receiver: self.receiver, + data: new_data, + type_map: Some(target_tm.clone()), + }) + } } fn fmt_data_value(val: &DataValue, f: &mut fmt::Formatter<'_>) -> fmt::Result { @@ -275,11 +320,26 @@ fn fmt_data_value(val: &DataValue, f: &mut fmt::Formatter<'_>) -> fmt::Result { } } +#[cfg(debug_assertions)] const BOLD_BLUE: &str = "\x1b[1;34m"; +#[cfg(not(debug_assertions))] +const BOLD_BLUE: &str = ""; +#[cfg(debug_assertions)] const GREEN: &str = "\x1b[32m"; +#[cfg(not(debug_assertions))] +const GREEN: &str = ""; +#[cfg(debug_assertions)] const YELLOW: &str = "\x1b[33m"; +#[cfg(not(debug_assertions))] +const YELLOW: &str = ""; +#[cfg(debug_assertions)] const ORANGE: &str = "\x1b[38;5;208m"; +#[cfg(not(debug_assertions))] +const ORANGE: &str = ""; +#[cfg(debug_assertions)] const RESET: &str = "\x1b[0m"; +#[cfg(not(debug_assertions))] +const RESET: &str = ""; impl fmt::Display for CommunicationValue { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { @@ -333,7 +393,7 @@ mod tests { #[test] fn test_flags_and_order_without_optional() { - let cv = CommunicationValue::from_comm_default(CommunicationType::ErrorParsing).with_id(0); + let cv = CommunicationValue::new(CommunicationType::ErrorParsing).with_id(0); let bytes = cv.to_bytes(); // [u32 len][u16 type][flags]... @@ -351,7 +411,7 @@ mod tests { #[test] fn test_flags_and_order_with_all_optional() { - let cv = CommunicationValue::from_comm_default(CommunicationType::ErrorBadVersion) + let cv = CommunicationValue::new(CommunicationType::ErrorBadVersion) .with_id(0xAABBCCDD) .with_sender(0x0000_1122_3344_5566) .with_receiver(0x0000_6677_8899_AABB); @@ -382,7 +442,7 @@ mod tests { #[test] fn test_roundtrip_complex() { - let cv = CommunicationValue::from_comm_default(CommunicationType::Disconnect) + let cv = CommunicationValue::new(CommunicationType::Disconnect) .with_id(1234) .with_sender(111) .with_receiver(222) @@ -415,6 +475,6 @@ mod tests { let mut bad = vec![0u8; 8]; // total_length claims more than available bad[0..4].copy_from_slice(&(1000u32.to_be_bytes())); - assert!(CommunicationValue::from_bytes(&bad).is_none()); + assert!(CommunicationValue::from_bytes(&bad).is_err()); } } diff --git a/codec/src/data_value.rs b/codec/src/data_value.rs index af1f361..bbe9fe4 100644 --- a/codec/src/data_value.rs +++ b/codec/src/data_value.rs @@ -311,7 +311,7 @@ impl DataValue { let kind = Self::kind_marker(value); buf.push(kind); - if kind == Self::KIND_BOOL_TRUE || kind == Self::KIND_BOOL_FALSE { + if kind == Self::KIND_BOOL_TRUE || kind == Self::KIND_BOOL_FALSE || kind == Self::KIND_NULL { let _ = buf.write_u16::(key.0); return true; } @@ -351,7 +351,7 @@ impl DataValue { let kind = Self::kind_marker(value); buf.push(kind); - if kind == Self::KIND_BOOL_TRUE || kind == Self::KIND_BOOL_FALSE { + if kind == Self::KIND_BOOL_TRUE || kind == Self::KIND_BOOL_FALSE || kind == Self::KIND_NULL { return true; } @@ -448,12 +448,14 @@ impl DataValue { for _ in 0..count { let kind = cursor.read_u8().ok()?; - if kind == Self::KIND_BOOL_TRUE || kind == Self::KIND_BOOL_FALSE { + if kind == Self::KIND_BOOL_TRUE || kind == Self::KIND_BOOL_FALSE || kind == Self::KIND_NULL { let key = DataTypeId(cursor.read_u16::().ok()?); let value = if kind == Self::KIND_BOOL_TRUE { DataValue::BoolTrue - } else { + } else if kind == Self::KIND_BOOL_FALSE { DataValue::BoolFalse + } else { + DataValue::Null }; entries.push((key, value)); continue; @@ -489,12 +491,15 @@ impl DataValue { for _ in 0..count { let kind = cursor.read_u8().ok()?; - if kind == Self::KIND_BOOL_TRUE { - out.push(DataValue::BoolTrue); - continue; - } - if kind == Self::KIND_BOOL_FALSE { - out.push(DataValue::BoolFalse); + if kind == Self::KIND_BOOL_TRUE || kind == Self::KIND_BOOL_FALSE || kind == Self::KIND_NULL { + let value = if kind == Self::KIND_BOOL_TRUE { + DataValue::BoolTrue + } else if kind == Self::KIND_BOOL_FALSE { + DataValue::BoolFalse + } else { + DataValue::Null + }; + out.push(value); continue; } diff --git a/codec/src/lib.rs b/codec/src/lib.rs index de8d858..b1e7074 100644 --- a/codec/src/lib.rs +++ b/codec/src/lib.rs @@ -1,25 +1,17 @@ pub mod communication_value; pub mod data_value; -pub mod util; pub use communication_value::CommunicationValue; pub use data_value::{DataKind, DataValue}; - -pub use util::rand_u32; +pub use mtp_common::CodecError; pub use mtp_type_map::{ communication_type_name, data_type_name, CommunicationType, CommunicationTypeId, DataType, - DataTypeId, PROTOCOL_VERSION, TypeMap, Version, + DataTypeId, TypeMap, Version, PROTOCOL_VERSION, }; -use mtp_common::CodecError; - -pub fn encode(_value: &DataValue, _typemap: &TypeMap) -> Result, CodecError> { - todo!() -} - -pub fn decode(_bytes: &[u8], _typemap: &TypeMap) -> Result { - todo!() +pub(crate) fn rand_u32() -> u32 { + rand::random() } #[cfg(feature = "registry")] diff --git a/codec/src/registry.rs b/codec/src/registry.rs index 40e97af..78757bf 100644 --- a/codec/src/registry.rs +++ b/codec/src/registry.rs @@ -1,8 +1,7 @@ -pub use registry::Registry; - -use crate::{CodecError, DataValue, decode, encode}; use mtp_type_map::Version; +pub use mtp_type_map::Registry; + /* * A version-aware codec that uses a multi-version registry to resolve * the correct TypeMap for encoding and decoding operations. @@ -17,22 +16,6 @@ impl VersionedCodec { Self { registry } } - pub fn encode(&self, value: &DataValue, version: Version) -> Result, CodecError> { - let typemap = self - .registry - .get(&version) - .ok_or(CodecError::UnknownVersion)?; - encode(value, typemap) - } - - pub fn decode(&self, bytes: &[u8], version: Version) -> Result { - let typemap = self - .registry - .get(&version) - .ok_or(CodecError::UnknownVersion)?; - decode(bytes, typemap) - } - pub fn negotiate(&self, client_versions: &[Version]) -> Option { self.registry.negotiate(client_versions) } diff --git a/codec/src/util.rs b/codec/src/util.rs deleted file mode 100644 index db90014..0000000 --- a/codec/src/util.rs +++ /dev/null @@ -1,9 +0,0 @@ -use rand::random; - -pub fn rand_u32() -> u32 { - random::() -} - -pub fn rand_u64() -> u64 { - random::() -} diff --git a/common/Cargo.toml b/common/Cargo.toml index e4e2bcf..ee28f47 100644 --- a/common/Cargo.toml +++ b/common/Cargo.toml @@ -10,8 +10,8 @@ wtransport = { version = "0.7.1", default-features = false, features = [ "quinn", "self-signed", ] } -rustls = { version = "0.23.40" } -quinn = { version = "0.11.9", default-features = false, features = [ +rustls = { version = "0.23.41" } +quinn = { version = "0.11.11", default-features = false, features = [ "rustls-aws-lc-rs", "rustls", ] } diff --git a/common/src/lib.rs b/common/src/lib.rs index 86471f9..339d426 100644 --- a/common/src/lib.rs +++ b/common/src/lib.rs @@ -1,16 +1,21 @@ use thiserror::Error; -pub enum RegistryError { - ReservedCommId(u16, String), -} -#[derive(Clone, Debug, PartialEq, Eq)] + +#[derive(Clone, Debug, Error, PartialEq, Eq)] pub enum CodecError { + #[error("Unknown version")] UnknownVersion, + #[error("Unknown communication type: {0}")] UnknownCommunicationType(String), + #[error("Unknown data type: {0}")] UnknownDataType(String), + #[error("Reserved communication type: {0}")] ReservedCommunicationType(u16), + #[error("Invalid encoding")] InvalidEncoding, + #[error("Crypto failed: {0}")] + CryptoFailed(String), } #[derive(Debug, Error, Clone)] @@ -39,17 +44,8 @@ pub enum CommunicationError { #[error("Loading Certificate error")] CertificateLoadFailed, - #[error("ParseBool error: {0}")] - ParseBool(#[from] std::str::ParseBoolError), - - #[error("ParseInt error: {0}")] - ParseInt(#[from] std::num::ParseIntError), - - #[error("ParseFloat error: {0}")] - ParseFloat(#[from] std::num::ParseFloatError), - - #[error("ParseAddr error: {0}")] - ParseAddr(#[from] std::net::AddrParseError), + #[error("Parse error: {0}")] + ParseError(String), #[error("Connection error: {0}")] ConnectionError(#[from] wtransport::error::ConnectionError), @@ -87,6 +83,9 @@ pub enum CommunicationError { #[error("Crypto Provider Install Error")] CryptoProviderInstallFailed, + #[error("Authentication failed: {0}")] + AuthenticationFailed(String), + #[error("Other: {0}")] Other(String), } diff --git a/crypto/Cargo.toml b/crypto/Cargo.toml index 1b37377..b00a2fa 100644 --- a/crypto/Cargo.toml +++ b/crypto/Cargo.toml @@ -21,6 +21,6 @@ serde = { version = "1", optional = true, features = ["derive"] } [features] default = ["chacha20poly1305", "ed25519-dalek", "hkdf", "sha2"] -full = ["chacha20poly1305", "aes-gcm", "ed25519-dalek", "hkdf", "sha2"] +full = ["default", "aes-gcm"] pqc = ["mlkem-tls", "ml-dsa"] serde = ["dep:serde"] diff --git a/crypto/src/helper.rs b/crypto/src/helper.rs index 6ca613b..beb75d7 100644 --- a/crypto/src/helper.rs +++ b/crypto/src/helper.rs @@ -57,15 +57,14 @@ impl MultiEncryptedMessage { pub fn from_bytes(bytes: &[u8]) -> Result { let mut offset = 0; let read_u16 = |off: &mut usize| -> Result { - let v = u16::from_be_bytes( - bytes - .get(*off..*off + 2) - .ok_or(CryptoError::DecryptionFailed)? - .try_into() - .unwrap(), - ); + let slice = bytes + .get(*off..*off + 2) + .ok_or(CryptoError::DecryptionFailed)?; + let arr: [u8; 2] = slice + .try_into() + .map_err(|_| CryptoError::DecryptionFailed)?; *off += 2; - Ok(v) + Ok(u16::from_be_bytes(arr)) }; let num = read_u16(&mut offset)? as usize; @@ -198,36 +197,4 @@ pub fn decrypt_multi( Err(CryptoError::DecryptionFailed) } -#[cfg(feature = "ed25519-dalek")] -pub fn verify_ed25519_sig( - public_key: &crate::keypair::SignaturePublicKey, - msg: &[u8], - signature: &[u8], -) -> Result<(), CryptoError> { - crate::sign::verify_ed25519(public_key, msg, signature) -} -#[cfg(feature = "ml-dsa")] -pub fn verify_ml_dsa_sig( - public_key: &crate::keypair::SignaturePqPublicKey, - msg: &[u8], - signature: &[u8], -) -> Result<(), CryptoError> { - crate::sign::verify_ml_dsa(public_key, msg, signature) -} - -/* - * Verify both an Ed25519 and ML-DSA signature (dual) against - * the public keys in a `PublicKeyBundle`. - */ -#[cfg(all(feature = "ed25519-dalek", feature = "ml-dsa"))] -pub fn verify_dual_sig( - public_keys: &crate::keypair::PublicKeyBundle, - msg: &[u8], - ed25519_sig: &[u8], - mldsa_sig: &[u8], -) -> Result<(), CryptoError> { - verify_ed25519_sig(&public_keys.sig_cl_public_key, msg, ed25519_sig)?; - verify_ml_dsa_sig(&public_keys.sig_pq_public_key, msg, mldsa_sig)?; - Ok(()) -} diff --git a/crypto/src/kem.rs b/crypto/src/kem.rs index 9e0c1c6..41c743c 100644 --- a/crypto/src/kem.rs +++ b/crypto/src/kem.rs @@ -12,7 +12,7 @@ pub struct HybridKem; #[cfg(feature = "mlkem-tls")] impl HybridKem { pub fn generate_keypair() -> (KemPrivateKey, KemPublicKey) { - let (dk, ek) = + let (ek, dk) = mlkem_tls::X25519MlKem768::keygen(&mut rand_core::OsRng); ( KemPrivateKey::new(dk.as_bytes().to_vec()), diff --git a/crypto/src/keypair.rs b/crypto/src/keypair.rs index f9cde57..e7d87a5 100644 --- a/crypto/src/keypair.rs +++ b/crypto/src/keypair.rs @@ -287,14 +287,18 @@ impl PublicKeyBundle { use crate::error::CryptoError; let mut offset = 0; - let kem_len = u16::from_be_bytes( - bytes - .get(offset..offset + 2) - .ok_or(CryptoError::InvalidKeyLength)? + let read_u16 = |off: &mut usize| -> Result { + let slice = bytes + .get(*off..*off + 2) + .ok_or(CryptoError::InvalidKeyLength)?; + let arr: [u8; 2] = slice .try_into() - .unwrap(), - ) as usize; - offset += 2; + .map_err(|_| CryptoError::InvalidKeyLength)?; + *off += 2; + Ok(u16::from_be_bytes(arr)) + }; + + let kem_len = read_u16(&mut offset)? as usize; let kem = KemPublicKey::new( bytes .get(offset..offset + kem_len) @@ -303,14 +307,7 @@ impl PublicKeyBundle { ); offset += kem_len; - let pq_len = u16::from_be_bytes( - bytes - .get(offset..offset + 2) - .ok_or(CryptoError::InvalidKeyLength)? - .try_into() - .unwrap(), - ) as usize; - offset += 2; + let pq_len = read_u16(&mut offset)? as usize; let pq = SignaturePqPublicKey::new( bytes .get(offset..offset + pq_len) @@ -319,14 +316,7 @@ impl PublicKeyBundle { ); offset += pq_len; - let cl_len = u16::from_be_bytes( - bytes - .get(offset..offset + 2) - .ok_or(CryptoError::InvalidKeyLength)? - .try_into() - .unwrap(), - ) as usize; - offset += 2; + let cl_len = read_u16(&mut offset)? as usize; let cl = SignaturePublicKey::new( bytes .get(offset..offset + cl_len) diff --git a/crypto/src/lib.rs b/crypto/src/lib.rs index 8bbc4f7..073d6e2 100644 --- a/crypto/src/lib.rs +++ b/crypto/src/lib.rs @@ -46,16 +46,248 @@ pub use hash::{sha256, sha256_double, Sha256Hasher}; pub use kdf::{derive_encryption_key, hkdf_expand, hkdf_extract}; #[cfg(feature = "mlkem-tls")] -pub use kem::HybridKem; +pub use kem::{Encapsulated, HybridKem}; #[cfg(all(feature = "mlkem-tls", feature = "chacha20poly1305", feature = "hkdf"))] pub use helper::{decrypt_multi, encrypt_multi, MultiEncryptedMessage, RecipientEntry}; -#[cfg(feature = "ed25519-dalek")] -pub use helper::verify_ed25519_sig; +/* ================================ TESTS ================================ */ +#[cfg(test)] +mod tests { + use super::*; -#[cfg(feature = "ml-dsa")] -pub use helper::verify_ml_dsa_sig; + #[cfg(feature = "chacha20poly1305")] + #[test] + fn aead_encrypt_decrypt() { + use crate::aead::{AeadDecrypt, AeadEncrypt, ChaCha20Poly1305}; + let key = [0xAB; 32]; + let cipher = ChaCha20Poly1305::new(key); + let ct = cipher.encrypt(b"hello world", b"aad").unwrap(); + let pt = cipher.decrypt(&ct, b"aad").unwrap(); + assert_eq!(pt, b"hello world"); + } -#[cfg(all(feature = "ed25519-dalek", feature = "ml-dsa"))] -pub use helper::verify_dual_sig; + #[cfg(feature = "chacha20poly1305")] + #[test] + fn aead_wrong_key_fails() { + use crate::aead::{AeadDecrypt, AeadEncrypt, ChaCha20Poly1305}; + let cipher_a = ChaCha20Poly1305::new([0xAB; 32]); + let cipher_b = ChaCha20Poly1305::new([0xCD; 32]); + let ct = cipher_a.encrypt(b"hello", b"").unwrap(); + assert!(cipher_b.decrypt(&ct, b"").is_err()); + } + + #[cfg(feature = "chacha20poly1305")] + #[test] + fn aead_wrong_aad_fails() { + use crate::aead::{AeadDecrypt, AeadEncrypt, ChaCha20Poly1305}; + let cipher = ChaCha20Poly1305::new([0xAB; 32]); + let ct = cipher.encrypt(b"hello", b"correct-aad").unwrap(); + assert!(cipher.decrypt(&ct, b"wrong-aad").is_err()); + } + + #[cfg(feature = "ed25519-dalek")] + #[test] + fn ed25519_sign_verify() { + let (signer, sk, pk) = Ed25519Signer::generate(); + let msg = b"test message"; + let sig = signer.sign(msg).unwrap(); + signer.verify(msg, &sig).unwrap(); + verify_ed25519(&pk, msg, &sig).unwrap(); + + let loaded = Ed25519Signer::new(&sk).unwrap(); + loaded.verify(msg, &sig).unwrap(); + } + + #[cfg(feature = "ed25519-dalek")] + #[test] + fn ed25519_wrong_sig_fails() { + let (signer, _, pk) = Ed25519Signer::generate(); + let msg = b"test message"; + let sig = signer.sign(msg).unwrap(); + assert!(verify_ed25519(&pk, b"wrong message", &sig).is_err()); + } + + #[cfg(feature = "ml-dsa")] + #[test] + fn mldsa_sign_verify() { + let (signer, sk, pk) = MlDsaSigner::generate(); + let msg = b"test message"; + let sig = signer.sign(msg).unwrap(); + signer.verify(msg, &sig).unwrap(); + verify_ml_dsa(&pk, msg, &sig).unwrap(); + + let loaded = MlDsaSigner::new(&sk, &pk).unwrap(); + loaded.verify(msg, &sig).unwrap(); + } + + #[cfg(feature = "ml-dsa")] + #[test] + fn mldsa_wrong_sig_fails() { + let (signer, _, pk) = MlDsaSigner::generate(); + let msg = b"test message"; + let sig = signer.sign(msg).unwrap(); + assert!(verify_ml_dsa(&pk, b"wrong message", &sig).is_err()); + } + + #[cfg(all(feature = "ed25519-dalek", feature = "ml-dsa"))] + #[test] + fn dual_sign_verify() { + use crate::sign::sign_dual; + + let (ed_signer, _, _) = Ed25519Signer::generate(); + let (ml_signer, _, _) = MlDsaSigner::generate(); + let dual = sign_dual(ed_signer.signing_key(), ml_signer.signing_key(), b"msg"); + dual + .verify( + ed_signer.verifying_key(), + ml_signer.verifying_key(), + b"msg", + ) + .unwrap(); + } + + #[cfg(all(feature = "ed25519-dalek", feature = "ml-dsa"))] + #[test] + fn dual_wrong_message_fails() { + use crate::sign::sign_dual; + + let (ed_signer, _, _) = Ed25519Signer::generate(); + let (ml_signer, _, _) = MlDsaSigner::generate(); + let dual = sign_dual(ed_signer.signing_key(), ml_signer.signing_key(), b"msg"); + assert!(dual + .verify(ed_signer.verifying_key(), ml_signer.verifying_key(), b"wrong") + .is_err()); + } + + #[cfg(feature = "hkdf")] + #[test] + fn hkdf_expand_produces_key() { + let key = derive_encryption_key(b"ikm", b"salt", b"context").unwrap(); + assert_eq!(key.len(), 32); + + let expanded = hkdf_expand(b"ikm", b"salt", b"info", 64).unwrap(); + assert_eq!(expanded.len(), 64); + } + + #[cfg(feature = "hkdf")] + #[test] + fn hkdf_different_info_different_key() { + let a = derive_encryption_key(b"ikm", b"salt", b"info-a").unwrap(); + let b = derive_encryption_key(b"ikm", b"salt", b"info-b").unwrap(); + assert_ne!(a, b); + } + + #[cfg(feature = "sha2")] + #[test] + fn sha256_hashes() { + let h = sha256(b"hello"); + assert_eq!(h.len(), 32); + + let h2 = sha256_double(b"hello"); + assert_eq!(h2.len(), 32); + assert_ne!(h, h2); + } + + #[cfg(feature = "sha2")] + #[test] + fn sha256_deterministic() { + assert_eq!(sha256(b"hello"), sha256(b"hello")); + } + + #[cfg(feature = "sha2")] + #[test] + fn sha256_hasher_incremental() { + let mut hasher = Sha256Hasher::new(); + hasher.update(b"hel"); + hasher.update(b"lo"); + let h = hasher.finalize(); + assert_eq!(h, sha256(b"hello")); + } + + #[test] + fn key_types_roundtrip() { + let data = vec![1u8, 2, 3, 4]; + let enc_pk = EncryptionPublicKey::new(data.clone()); + assert_eq!(enc_pk.as_bytes(), &data); + let enc_sk = EncryptionPrivateKey::new(data.clone()); + assert_eq!(enc_sk.as_bytes(), &data); + + let sig_pk = SignaturePublicKey::new(data.clone()); + assert_eq!(sig_pk.as_bytes(), &data); + let sig_sk = SignaturePrivateKey::new(data.clone()); + assert_eq!(sig_sk.as_bytes(), &data); + + let kem_pk = KemPublicKey::new(data.clone()); + assert_eq!(kem_pk.as_bytes(), &data); + let kem_sk = KemPrivateKey::new(data.clone()); + assert_eq!(kem_sk.as_bytes(), &data); + + let sig_pq_pk = SignaturePqPublicKey::new(data.clone()); + assert_eq!(sig_pq_pk.as_bytes(), &data); + let sig_pq_sk = SignaturePqPrivateKey::new(data.clone()); + assert_eq!(sig_pq_sk.as_bytes(), &data); + } + + #[cfg(all(feature = "mlkem-tls", feature = "ml-dsa", feature = "ed25519-dalek"))] + #[test] + fn keyring_generate_and_bundle() { + let kr = Keyring::generate(); + let bundle = kr.public_key_bundle(); + assert!(!bundle.kem_public_key.as_bytes().is_empty()); + assert!(!bundle.sig_pq_public_key.as_bytes().is_empty()); + assert!(!bundle.sig_cl_public_key.as_bytes().is_empty()); + } + + #[cfg(all(feature = "mlkem-tls", feature = "ml-dsa", feature = "ed25519-dalek"))] + #[test] + fn keyring_serialize_roundtrip() { + let kr = Keyring::generate(); + let bytes = kr.to_bytes(); + let loaded = Keyring::from_bytes(&bytes).unwrap(); + assert_eq!(kr.kem_public_key.as_bytes(), loaded.kem_public_key.as_bytes()); + assert_eq!(kr.sig_pq_public_key.as_bytes(), loaded.sig_pq_public_key.as_bytes()); + assert_eq!(kr.sig_cl_public_key.as_bytes(), loaded.sig_cl_public_key.as_bytes()); + } + + #[cfg(all(feature = "mlkem-tls", feature = "ml-dsa", feature = "ed25519-dalek"))] + #[test] + fn public_key_bundle_serialize_roundtrip() { + let kr = Keyring::generate(); + let bundle = kr.public_key_bundle(); + let bytes = bundle.as_bytes(); + let loaded = PublicKeyBundle::from_bytes(&bytes).unwrap(); + assert_eq!(bundle.kem_public_key.as_bytes(), loaded.kem_public_key.as_bytes()); + assert_eq!( + bundle.sig_pq_public_key.as_bytes(), + loaded.sig_pq_public_key.as_bytes() + ); + assert_eq!( + bundle.sig_cl_public_key.as_bytes(), + loaded.sig_cl_public_key.as_bytes() + ); + } + + #[cfg(feature = "mlkem-tls")] + #[test] + fn hybrid_kem_roundtrip() { + let (sk, pk) = HybridKem::generate_keypair(); + let enc = HybridKem::encapsulate(&pk).unwrap(); + let ss = HybridKem::decapsulate(&sk, &enc.ciphertext).unwrap(); + assert_eq!(enc.shared_secret, ss); + } + + #[cfg(all(feature = "mlkem-tls", feature = "chacha20poly1305", feature = "hkdf"))] + #[test] + fn encrypt_multi_roundtrip() { + use crate::keypair::Keyring; + use crate::helper::{encrypt_multi, decrypt_multi}; + + let kr = Keyring::generate(); + let entities = vec![kr.public_key_bundle()]; + let msg = b"secret data"; + let ct = encrypt_multi(msg, b"aad", &entities).unwrap(); + let pt = decrypt_multi(&ct, b"aad", &kr).unwrap(); + assert_eq!(pt, msg); + } +} diff --git a/crypto/src/sign.rs b/crypto/src/sign.rs index fbc641a..aeb593f 100644 --- a/crypto/src/sign.rs +++ b/crypto/src/sign.rs @@ -46,6 +46,14 @@ impl Ed25519Signer { pub fn public_key(&self) -> SignaturePublicKey { SignaturePublicKey::new(self.public.to_bytes().to_vec()) } + + pub fn signing_key(&self) -> &ed25519_dalek::SigningKey { + &self.secret + } + + pub fn verifying_key(&self) -> &ed25519_dalek::VerifyingKey { + &self.public + } } #[cfg(feature = "ed25519-dalek")] @@ -104,23 +112,26 @@ impl MlDsaSigner { secret_key: &SignaturePqPrivateKey, public_key: &SignaturePqPublicKey, ) -> Result { - let encoded_sk = - ml_dsa::EncodedSigningKey::::try_from(secret_key.as_bytes()) - .map_err(|_| CryptoError::KeyGenerationFailed)?; - let secret = ml_dsa::SigningKey::::decode(&encoded_sk); + let seed_bytes: [u8; 32] = secret_key + .as_bytes() + .try_into() + .map_err(|_| CryptoError::KeyGenerationFailed)?; + let seed = ml_dsa::Seed::from(seed_bytes); + let secret = ml_dsa::SigningKey::::from_seed(&seed); + let encoded_pk = ml_dsa::EncodedVerifyingKey::::try_from(public_key.as_bytes()) .map_err(|_| CryptoError::KeyGenerationFailed)?; let public = ml_dsa::VerifyingKey::::decode(&encoded_pk); + Ok(Self { secret, public }) } pub fn generate() -> (Self, SignaturePqPrivateKey, SignaturePqPublicKey) { - use ml_dsa::KeyGen; - let kp = ml_dsa::MlDsa65::key_gen(&mut rand_core::OsRng); - let secret = kp.signing_key().clone(); - let public = kp.verifying_key().clone(); - let priv_key = SignaturePqPrivateKey::new(secret.encode().to_vec()); + use ml_dsa::{Generate, Keypair}; + let secret = ml_dsa::SigningKey::::generate(); + let public = secret.verifying_key(); + let priv_key = SignaturePqPrivateKey::new(secret.to_seed().to_vec()); let pub_key = SignaturePqPublicKey::new(public.encode().to_vec()); let signer = Self { secret, public }; (signer, priv_key, pub_key) @@ -142,20 +153,17 @@ impl MlDsaSigner { #[cfg(feature = "ml-dsa")] impl SignatureScheme for MlDsaSigner { fn sign(&self, msg: &[u8]) -> Result, CryptoError> { - use ml_dsa::signature::Signer; - let signature = self.secret.sign(msg); + use ml_dsa::Signer; + let signature = self.secret.try_sign(msg) + .map_err(|_| CryptoError::SigningFailed)?; Ok(signature.encode().to_vec()) } fn verify(&self, msg: &[u8], signature: &[u8]) -> Result<(), CryptoError> { - use ml_dsa::signature::Verifier; - let encoded_sig = - ml_dsa::EncodedSignature::::try_from(signature) - .map_err(|_| CryptoError::InvalidSignature)?; - let sig = ml_dsa::Signature::::decode(&encoded_sig) - .ok_or(CryptoError::InvalidSignature)?; - self.public - .verify(msg, &sig) + use ml_dsa::Verifier; + let sig = ml_dsa::Signature::::try_from(signature) + .map_err(|_| CryptoError::InvalidSignature)?; + self.public.verify(msg, &sig) .map_err(|_| CryptoError::VerificationFailed) } } @@ -166,23 +174,22 @@ pub fn verify_ml_dsa( msg: &[u8], signature: &[u8], ) -> Result<(), CryptoError> { - use ml_dsa::signature::Verifier; + use ml_dsa::Verifier; let encoded_pk = ml_dsa::EncodedVerifyingKey::::try_from(public_key.as_bytes()) .map_err(|_| CryptoError::InvalidSignature)?; let public = ml_dsa::VerifyingKey::::decode(&encoded_pk); - let encoded_sig = - ml_dsa::EncodedSignature::::try_from(signature) - .map_err(|_| CryptoError::InvalidSignature)?; - let sig = ml_dsa::Signature::::decode(&encoded_sig) - .ok_or(CryptoError::InvalidSignature)?; + + let sig = ml_dsa::Signature::::try_from(signature) + .map_err(|_| CryptoError::InvalidSignature)?; public .verify(msg, &sig) .map_err(|_| CryptoError::VerificationFailed) } +#[derive(Clone, Debug, PartialEq, Eq)] pub struct DualSignature { pub ed25519: Vec, pub mldsa: Vec, @@ -194,12 +201,18 @@ pub fn sign_dual( mldsa_sk: &ml_dsa::SigningKey, message: &[u8], ) -> DualSignature { - use ed25519_dalek::Signer; - - DualSignature { - ed25519: ed25519_sk.sign(message).to_bytes().to_vec(), - mldsa: mldsa_sk.sign(message).encode().to_vec(), - } + let ed25519 = { + use ed25519_dalek::Signer; + ed25519_sk.sign(message).to_bytes().to_vec() + }; + let mldsa = { + use ml_dsa::Signer; + mldsa_sk.try_sign(message) + .expect("ML-DSA signing failed") + .encode() + .to_vec() + }; + DualSignature { ed25519, mldsa } } impl DualSignature { @@ -210,22 +223,23 @@ impl DualSignature { mldsa_vk: &ml_dsa::VerifyingKey, message: &[u8], ) -> Result<(), CryptoError> { - use ed25519_dalek::Verifier; - let ed_sig = ed25519_dalek::Signature::from_slice(&self.ed25519) .map_err(|_| CryptoError::InvalidSignature)?; - ed25519_vk - .verify(message, &ed_sig) - .map_err(|_| CryptoError::VerificationFailed)?; + { + use ed25519_dalek::Verifier; + ed25519_vk + .verify(message, &ed_sig) + .map_err(|_| CryptoError::VerificationFailed)?; + } - let encoded_sig = - ml_dsa::EncodedSignature::::try_from(self.mldsa.as_slice()) - .map_err(|_| CryptoError::InvalidSignature)?; - let ml_sig = ml_dsa::Signature::::decode(&encoded_sig) - .ok_or(CryptoError::InvalidSignature)?; - mldsa_vk - .verify(message, &ml_sig) - .map_err(|_| CryptoError::VerificationFailed)?; + let ml_sig = ml_dsa::Signature::::try_from(self.mldsa.as_slice()) + .map_err(|_| CryptoError::InvalidSignature)?; + { + use ml_dsa::Verifier; + mldsa_vk + .verify(message, &ml_sig) + .map_err(|_| CryptoError::VerificationFailed)?; + } Ok(()) } diff --git a/example-type-maps.yaml b/example-type-maps.yaml index 0a78ca6..0041f48 100644 --- a/example-type-maps.yaml +++ b/example-type-maps.yaml @@ -33,9 +33,11 @@ protocol_version: "0.0" # Description: 4 # Timestamp: 5 # Id: 6 -# Nonce: 7 -# PublicKeys: 8 -# Signature: 9 +# ClientNonce: 7 +# ServerNonce: 8 +# PublicKeys: 9 +# Signature: 10 +# Connected: 11 # # If a Type can't be used it will be mapped to 0 diff --git a/example-usage/.gitignore b/example-usage/.gitignore new file mode 100644 index 0000000..cf5e6e3 --- /dev/null +++ b/example-usage/.gitignore @@ -0,0 +1,4 @@ +client_keys.json +host_keys.json +host_sig_pk.bin +clients.json diff --git a/example-usage/client/Cargo.toml b/example-usage/client/Cargo.toml index 67a0fed..1f2a829 100644 --- a/example-usage/client/Cargo.toml +++ b/example-usage/client/Cargo.toml @@ -8,6 +8,9 @@ name = "mtp-example-client" path = "src/main.rs" [dependencies] -mtp-client = { path = "../../client" } +mtp-client = { path = "../../client", features = ["crypto"] } mtp-codec = { path = "../../codec" } +mtp-crypto = { path = "../../crypto" } tokio = { version = "1", features = ["full"] } +serde_json = "1" +hex = "0.4" diff --git a/example-usage/client/src/main.rs b/example-usage/client/src/main.rs index d9be8a2..88ed12e 100644 --- a/example-usage/client/src/main.rs +++ b/example-usage/client/src/main.rs @@ -1,23 +1,77 @@ use std::fs; -use mtp_client::{ClientConfig, MTPClient}; +use mtp_client::{ClientConfig, MTPClient, MTPConnection}; use mtp_codec::{CommunicationType, CommunicationValue, DataType, DataValue}; +use mtp_crypto::{ + Ed25519Signer, EncryptionPrivateKey, EncryptionPublicKey, KeyGroup, Keyring, + KemPrivateKey, KemPublicKey, SignaturePqPrivateKey, SignaturePqPublicKey, + SignaturePrivateKey, SignaturePublicKey, +}; -#[tokio::main] -async fn main() -> Result<(), Box> { - let cert_pem = fs::read("server.pem") - .expect("Missing server.pem server generates it"); +/* + * Either load a saved keyring from disk (login) or generate a fresh + * keypair and register with the server. Returns the established + * connection and the local keyring. + */ +async fn connect_or_register( + config: ClientConfig, + host_public_key: SignaturePublicKey, + client_key_path: &str, +) -> Result<(MTPConnection, Keyring), Box> { + if let Ok(data) = fs::read_to_string(client_key_path) { + let json: serde_json::Value = serde_json::from_str(&data)?; + let client_id = json["client_id"].as_u64().expect("Invalid client_id"); + let keyring = Keyring::from_bytes( + &hex::decode(json["keyring"].as_str().expect("Missing keyring"))?, + )?; - println!("Connecting to 127.0.0.1:8080 ..."); + println!("Loaded client keys (ID: {})", client_id); - let config = ClientConfig { - url: "https://127.0.0.1:8080".into(), - server_cert: Some(cert_pem), - }; + let keygroup = KeyGroup::new( + EncryptionPublicKey::new(vec![]), + EncryptionPrivateKey::new(vec![]), + keyring.sig_cl_public_key.clone(), + SignaturePrivateKey::new(keyring.sig_cl_secret_key.as_bytes().to_vec()), + ); - let conn = MTPClient::connect(config, 1001).await?; - println!("Connected (version {})", conn.version); + let conn = MTPClient::auth_connect(config, client_id, keygroup, host_public_key).await?; + println!("Authenticated (version {})", conn.version); + Ok((conn, keyring)) + } else { + println!("No existing keys found - registering new client"); + let (_signer, sig_sk, sig_pk) = Ed25519Signer::generate(); + let keyring = Keyring::new( + KemPublicKey::new(vec![]), + KemPrivateKey::new(vec![]), + SignaturePqPublicKey::new(vec![]), + SignaturePqPrivateKey::new(vec![]), + sig_pk, + sig_sk, + ); + + // Serialise before the move so we can persist and return the keyring + let keyring_bytes = keyring.to_bytes(); + let conn = MTPClient::auth_register(config, keyring, host_public_key).await?; + println!("Registered with ID: {}", conn.client_id); + + let json = serde_json::json!({ + "client_id": conn.client_id, + "keyring": hex::encode(&keyring_bytes), + }); + fs::write(client_key_path, serde_json::to_string_pretty(&json)?)?; + println!("Saved client keys -> {client_key_path}"); + + let keyring = Keyring::from_bytes(&keyring_bytes)?; + Ok((conn, keyring)) + } +} + +/* + * Send a Ping greeting with description, timestamp, and data fields, + * then print the server's Pong response. + */ +async fn send_ping_and_receive(conn: &MTPConnection) -> Result<(), Box> { let timestamp = std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) .unwrap() @@ -27,22 +81,37 @@ async fn main() -> Result<(), Box> { .add_typed_default(DataType::Description, DataValue::Str("Hello MTP!".into())) .add_typed_default(DataType::Timestamp, DataValue::UnsignedNumber(timestamp as u128)) .add_typed_default(DataType::Data, DataValue::UnsignedNumber(42)) - .with_sender(1001); + .with_sender(conn.client_id); println!("Sending: {greeting}"); conn.sender.send(&greeting).await?; match conn.receiver.receive().await { - Ok(msg) => { - println!("Received: {msg}"); - } - Err(e) => { - eprintln!("Receive error: {e}"); - } + Ok(msg) => println!("Received: {msg}"), + Err(e) => eprintln!("Receive error: {e}"), } conn.sender.close(); + Ok(()) +} + +#[tokio::main] +async fn main() -> Result<(), Box> { + let cert_pem = fs::read("server.pem").expect("Missing server.pem - run server first"); + let host_public_key = SignaturePublicKey::new( + fs::read("host_sig_pk.bin").expect("Missing host_sig_pk.bin - run server first"), + ); + + println!("Connecting to 127.0.0.1:8080 ..."); + + let config = ClientConfig { + url: "https://127.0.0.1:8080".into(), + server_cert: Some(cert_pem), + }; + + let (conn, _keyring) = connect_or_register(config, host_public_key, "client_keys.json").await?; + send_ping_and_receive(&conn).await?; + println!("\nDone"); - Ok(()) } diff --git a/example-usage/server/Cargo.toml b/example-usage/server/Cargo.toml index 6da0426..19f3031 100644 --- a/example-usage/server/Cargo.toml +++ b/example-usage/server/Cargo.toml @@ -8,7 +8,10 @@ name = "mtp-example-server" path = "src/main.rs" [dependencies] -mtp-host = { path = "../../host" } +mtp-host = { path = "../../host", features = ["crypto"] } mtp-codec = { path = "../../codec" } +mtp-crypto = { path = "../../crypto", features = ["serde"] } rcgen = "0.14" tokio = { version = "1", features = ["full"] } +serde_json = "1" +hex = "0.4" diff --git a/example-usage/server/src/main.rs b/example-usage/server/src/main.rs index e58eaae..b8aea9c 100644 --- a/example-usage/server/src/main.rs +++ b/example-usage/server/src/main.rs @@ -1,79 +1,188 @@ +use std::collections::HashMap; use std::fs; +use std::sync::{Arc, Mutex}; -use mtp_codec::{CommunicationType, CommunicationValue, DataType, DataValue, TypeMap}; +use mtp_codec::{CommunicationType, CommunicationValue, DataType, DataTypeId, DataValue, TypeMap}; use mtp_host::{HostConfig, MTPHost}; +use mtp_crypto::{ + Ed25519Signer, KemPrivateKey, KemPublicKey, Keyring, PublicKeyBundle, + SignaturePqPrivateKey, SignaturePqPublicKey, +}; + +/* + * Load an existing TLS certificate and key pair from disk. If neither + * file exists, generate a self-signed certificate so the server is + * immediately usable without external tooling. + */ +fn load_or_generate_tls(cert_path: &str, key_path: &str) -> Result<(Vec, Vec), Box> { + if let (Ok(c), Ok(k)) = (fs::read(cert_path), fs::read(key_path)) { + println!("Using existing TLS cert from {cert_path}"); + return Ok((c, k)); + } + + println!("Generating self-signed TLS certificate ..."); + let key_pair = rcgen::KeyPair::generate()?; + let params = rcgen::CertificateParams::new(vec!["localhost".into(), "127.0.0.1".into()])?; + let cert = params.self_signed(&key_pair)?; + + let cert_str = cert.pem(); + let key_str = key_pair.serialize_pem(); + + fs::write(cert_path, cert_str.as_bytes())?; + fs::write(key_path, key_str.as_bytes())?; + println!("Wrote {cert_path} and {key_path}"); + + Ok((cert_str.into_bytes(), key_str.into_bytes())) +} + +/* + * Load the host's Ed25519 identity from a JSON file, or generate a + * fresh one and persist it. Clients need the corresponding public key + * (exported separately as host_sig_pk.bin) to authenticate the host + * during the handshake. + */ +fn load_or_generate_host_keys(path: &str) -> Result<(u64, Keyring), Box> { + if let Ok(data) = fs::read_to_string(path) { + let json: serde_json::Value = serde_json::from_str(&data)?; + let hid = json["host_id"].as_u64().unwrap_or(1); + let keyring = Keyring::from_bytes(&hex::decode(json["keyring"].as_str().unwrap())?)?; + println!("Loaded host keys (ID: {})", hid); + return Ok((hid, keyring)); + } + + let (_signer, sig_sk, sig_pk) = Ed25519Signer::generate(); + let keyring = Keyring::new( + KemPublicKey::new(vec![]), + KemPrivateKey::new(vec![]), + SignaturePqPublicKey::new(vec![]), + SignaturePqPrivateKey::new(vec![]), + sig_pk, + sig_sk, + ); + + let json = serde_json::json!({ + "host_id": 1, + "keyring": hex::encode(keyring.to_bytes()), + }); + fs::write(path, serde_json::to_string_pretty(&json)?)?; + println!("Generated host keys -> {path}"); + Ok((1u64, keyring)) +} + +/* + * Load the client database from disk. Each entry maps a numeric + * client ID to its PublicKeyBundle. next_id starts one past the + * highest known ID (or 1000 if the DB is empty). + */ +fn load_client_db( + path: &str, +) -> Result< + ( + Arc>>, + Arc>, + ), + Box, +> { + let clients: Arc>> = Arc::new(Mutex::new( + if let Ok(data) = fs::read_to_string(path) { + serde_json::from_str(&data).unwrap_or_default() + } else { + HashMap::new() + }, + )); + let next_id = Arc::new(Mutex::new( + clients.lock().unwrap().keys().max().unwrap_or(&999) + 1, + )); + Ok((clients, next_id)) +} + +/* + * Build a Pong response carrying a description, a Unix timestamp, and + * a custom payload according to the negotiated type map. + */ +fn build_pong_response(tm: &TypeMap) -> CommunicationValue { + let timestamp = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_secs(); + + CommunicationValue::from_comm(CommunicationType::Pong, tm) + .add_data( + DataTypeId(tm.data_id_enum(DataType::Description).unwrap()), + DataValue::Str("Hello from server!".into()), + ) + .add_data( + DataTypeId(tm.data_id_enum(DataType::Timestamp).unwrap()), + DataValue::UnsignedNumber(timestamp as u128), + ) + .add_data( + DataTypeId(tm.data_id_enum(DataType::Data).unwrap()), + DataValue::Str("custom payload".into()), + ) +} #[tokio::main] async fn main() -> Result<(), Box> { - let cert_path = "server.pem"; - let key_path = "server.key"; + let (cert_pem, key_pem) = load_or_generate_tls("server.pem", "server.key")?; + let (host_id, host_keyring) = load_or_generate_host_keys("host_keys.json")?; - let (cert_pem, key_pem) = if let (Ok(c), Ok(k)) = (fs::read(cert_path), fs::read(key_path)) { - println!("Using existing TLS cert from {cert_path}"); - (c, k) - } else { - println!("Generating self-signed TLS certificate ..."); + // Export the host's Ed25519 public key so clients can verify it + fs::write( + "host_sig_pk.bin", + host_keyring.sig_cl_public_key.as_bytes(), + )?; - let key_pair = rcgen::KeyPair::generate()?; - let params = rcgen::CertificateParams::new(vec![ - "localhost".into(), - "127.0.0.1".into(), - ])?; - let cert = params.self_signed(&key_pair)?; + let (clients, next_id) = load_client_db("clients.json")?; - let cert_str = cert.pem(); - let key_str = key_pair.serialize_pem(); + // Clone the Arc so each closure owns its own reference + let clients_for_get = clients.clone(); + let get_existing_user = Box::new(move |id: u64| -> Option { + clients_for_get.lock().unwrap().get(&id).cloned() + }); - fs::write(cert_path, cert_str.as_bytes())?; - fs::write(key_path, key_str.as_bytes())?; - println!("Wrote {cert_path} and {key_path}"); - - (cert_str.into_bytes(), key_str.into_bytes()) - }; + let clients_for_register = clients.clone(); + let next_id_for_register = next_id.clone(); + let clients_path = "clients.json".to_string(); + let complete_register = Box::new(move |bundle: PublicKeyBundle| -> u64 { + let mut db = clients_for_register.lock().unwrap(); + let mut nid = next_id_for_register.lock().unwrap(); + let id = *nid; + *nid += 1; + db.insert(id, bundle); + fs::write(&clients_path, serde_json::to_string_pretty(&*db).unwrap()).ok(); + println!("Registered new client with ID: {}", id); + id + }); println!("Starting MTP server on port 8080 ..."); let config = HostConfig { - ip: "::".into(), port: 8080, tls_fullchain: cert_pem, tls_key: key_pem, + require_authentication: true, + host_id, + host_keyring, + get_existing_user, + complete_register, }; let mut host = MTPHost::new(config).await?; println!("Server listening on {}", host.local_addr()); while let Some(conn) = host.accept().await { - println!("\n--- New connection accepted (version {}) ---", conn.version); + println!( + "\n--- New authenticated connection (version {}) ---", + conn.version + ); + println!("Client ID: {}", conn.client_id); let tm: &TypeMap = conn.codec.registry().get(&conn.version).unwrap(); match conn.receiver.receive().await { Ok(msg) => { println!("Received: {msg}"); - - let timestamp = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs(); - - let response = CommunicationValue::from_comm( - CommunicationType::Pong, - tm, - ) - .add_data( - mtp_codec::DataTypeId(tm.data_id_enum(DataType::Description).unwrap()), - DataValue::Str("Hello from server!".into()), - ) - .add_data( - mtp_codec::DataTypeId(tm.data_id_enum(DataType::Timestamp).unwrap()), - DataValue::UnsignedNumber(timestamp as u128), - ) - .add_data( - mtp_codec::DataTypeId(tm.data_id_enum(DataType::Data).unwrap()), - DataValue::Str("custom payload".into()), - ); - + let response = build_pong_response(tm); println!("Sending: {response}"); conn.sender.send(&response).await?; } diff --git a/host/Cargo.toml b/host/Cargo.toml index fd26447..482b9f4 100644 --- a/host/Cargo.toml +++ b/host/Cargo.toml @@ -7,4 +7,9 @@ edition = "2024" mtp-common = { path = "../common" } mtp-codec = { path = "../codec", features = ["registry"] } mtp-transport = { path = "../transport", features = ["host"] } +mtp-crypto = { path = "../crypto", optional = true } tokio = { version = "1", features = ["full"] } +rand = "0.8" + +[features] +crypto = ["dep:mtp-crypto", "mtp-codec/crypto"] diff --git a/host/src/lib.rs b/host/src/lib.rs index 739bef0..6d6b24e 100644 --- a/host/src/lib.rs +++ b/host/src/lib.rs @@ -1,3 +1,5 @@ +#[cfg(feature = "crypto")] +use mtp_codec::DataType; use mtp_codec::{ CommunicationValue, DataTypeId, DataValue, Version, registry::{Registry, VersionedCodec}, @@ -7,10 +9,28 @@ use mtp_transport::{Policy, Receiver, Sender}; // Host configuration. pub struct HostConfig { - pub ip: String, pub port: u16, pub tls_fullchain: Vec, pub tls_key: Vec, + + #[cfg(feature = "crypto")] + pub require_authentication: bool, + #[cfg(feature = "crypto")] + pub host_id: u64, + #[cfg(feature = "crypto")] + pub host_keyring: mtp_crypto::Keyring, + #[cfg(feature = "crypto")] + pub get_existing_user: Box Option + Send>, + #[cfg(feature = "crypto")] + pub complete_register: Box u64 + Send>, +} + +#[cfg(feature = "crypto")] +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum AuthState { + Pending, + Authenticated, + Failed, } // A connection that has completed version negotiation. @@ -19,12 +39,20 @@ pub struct MTPConnection { pub codec: VersionedCodec, pub sender: Sender, pub receiver: Receiver, + #[cfg(feature = "crypto")] + pub auth_state: AuthState, + #[cfg(feature = "crypto")] + pub client_id: u64, + #[cfg(feature = "crypto")] + pub client_public_key: Option, } // High-level MTP host with built-in version negotiation. pub struct MTPHost { transport: mtp_transport::Host, registry: Registry, + #[cfg(feature = "crypto")] + config: HostConfig, } impl MTPHost { @@ -33,8 +61,8 @@ impl MTPHost { let transport = mtp_transport::host( config.port, - config.tls_fullchain, - config.tls_key, + config.tls_fullchain.clone(), + config.tls_key.clone(), Policy::default(), ) .await?; @@ -42,6 +70,8 @@ impl MTPHost { Ok(Self { transport, registry, + #[cfg(feature = "crypto")] + config, }) } @@ -55,6 +85,11 @@ impl MTPHost { pub async fn accept(&mut self) -> Option { let (sender, receiver) = self.transport.next().await?; + #[cfg(feature = "crypto")] + if self.config.require_authentication { + return self.accept_authenticated(sender, receiver).await; + } + /* Read the first message (always encoded with reserved types). */ let first_msg = receiver.receive().await.ok()?; @@ -76,6 +111,12 @@ impl MTPHost { codec, sender, receiver, + #[cfg(feature = "crypto")] + auth_state: AuthState::Authenticated, + #[cfg(feature = "crypto")] + client_id: 0, + #[cfg(feature = "crypto")] + client_public_key: None, }) } @@ -88,6 +129,193 @@ impl MTPHost { } } +#[cfg(feature = "crypto")] +impl MTPHost { + async fn accept_authenticated( + &mut self, + sender: Sender, + receiver: Receiver, + ) -> Option { + use mtp_crypto::{Ed25519Signer, PublicKeyBundle, SignatureScheme, verify_ed25519}; + + // 1. Send host greeting + let host_nonce: u128 = rand::random(); + let host_bundle = self.config.host_keyring.public_key_bundle(); + + let greeting = + CommunicationValue::new(mtp_codec::CommunicationType::IdentificationResponse) + .add_typed_default( + DataType::Id, + DataValue::UnsignedNumber(self.config.host_id as u128), + ) + .add_typed_default( + DataType::PublicKeys, + DataValue::Bytes(host_bundle.as_bytes()), + ) + .add_typed_default(DataType::ServerNonce, DataValue::UnsignedNumber(host_nonce)); + + sender.send(&greeting).await.ok()?; + + // 2. Receive client message + let msg = receiver.receive().await.ok()?; + + // Extract common fields + let version_str = match msg.get_data(DataTypeId(3)) { + DataValue::Str(s) => s.clone(), + _ => { + sender.close(); + return None; + } + }; + + let client_version = Version::parse(&version_str)?; + + let client_nonce = match msg.get_data(DataTypeId(7)) { + DataValue::UnsignedNumber(n) => *n, + _ => { + sender.close(); + return None; + } + }; + + let sig_bytes = match msg.get_data(DataTypeId(10)) { + DataValue::Bytes(b) => b.clone(), + _ => { + sender.close(); + return None; + } + }; + + let (assigned_id, client_bundle) = if msg.get_type() == mtp_codec::CommunicationTypeId(15) { + // LOGIN + let cid = match msg.get_data(DataTypeId(6)) { + DataValue::UnsignedNumber(n) => *n as u64, + _ => { + sender.close(); + return None; + } + }; + + let bundle = match (self.config.get_existing_user)(cid) { + Some(b) => b, + None => { + let rejection = CommunicationValue::new( + mtp_codec::CommunicationType::IdentificationResponse, + ) + .add_typed_default(DataType::Connected, DataValue::BoolFalse) + .add_typed_default( + DataType::ClientNonce, + DataValue::UnsignedNumber(client_nonce), + ); + let _ = sender.send(&rejection).await; + sender.close(); + return None; + } + }; + + let mut sig_payload = Vec::new(); + sig_payload.extend_from_slice(version_str.as_bytes()); + sig_payload.extend_from_slice(&cid.to_be_bytes()); + sig_payload.extend_from_slice(&client_nonce.to_be_bytes()); + + /* ===== Signature ===== */ + if verify_ed25519(&bundle.sig_cl_public_key, &sig_payload, &sig_bytes).is_err() { + let rejection = + CommunicationValue::new(mtp_codec::CommunicationType::IdentificationResponse) + .add_typed_default(DataType::Connected, DataValue::BoolFalse) + .add_typed_default( + DataType::ClientNonce, + DataValue::UnsignedNumber(client_nonce), + ); + let _ = sender.send(&rejection).await; + sender.close(); + return None; + } + /* ===== End Signature ===== */ + + (cid, bundle) + } else if msg.get_type() == mtp_codec::CommunicationTypeId(14) { + // REGISTER + let bundle = match msg.get_data(DataTypeId(8)) { + DataValue::Bytes(b) => PublicKeyBundle::from_bytes(b).ok()?, + _ => { + sender.close(); + return None; + } + }; + let pk_bytes = bundle.as_bytes(); + + let mut sig_payload = Vec::new(); + sig_payload.extend_from_slice(version_str.as_bytes()); + sig_payload.extend_from_slice(&client_nonce.to_be_bytes()); + sig_payload.extend_from_slice(&pk_bytes); + + /* ===== Signature ===== */ + if verify_ed25519(&bundle.sig_cl_public_key, &sig_payload, &sig_bytes).is_err() { + let rejection = + CommunicationValue::new(mtp_codec::CommunicationType::RegisterResponse) + .add_typed_default(DataType::Connected, DataValue::BoolFalse) + .add_typed_default( + DataType::ClientNonce, + DataValue::UnsignedNumber(client_nonce), + ); + let _ = sender.send(&rejection).await; + sender.close(); + return None; + } + /* ===== End Signature ===== */ + + let new_id = (self.config.complete_register)(bundle.clone()); + (new_id, bundle) + } else { + sender.close(); + return None; + }; + + // 4. Send success response + let new_nonce: u128 = rand::random(); + + let mut host_sig_payload = Vec::new(); + host_sig_payload.extend_from_slice(b"true"); + host_sig_payload.extend_from_slice(&assigned_id.to_be_bytes()); + host_sig_payload.extend_from_slice(&client_nonce.to_be_bytes()); + host_sig_payload.extend_from_slice(&new_nonce.to_be_bytes()); + + let host_signer = Ed25519Signer::new(&self.config.host_keyring.sig_cl_secret_key).ok()?; + + /* ===== Signature ===== */ + let host_sig = host_signer.sign(&host_sig_payload).ok()?; + /* ===== End Signature ===== */ + + let response = + CommunicationValue::new(mtp_codec::CommunicationType::IdentificationResponse) + .add_typed_default(DataType::Connected, DataValue::BoolTrue) + .add_typed_default( + DataType::ClientNonce, + DataValue::UnsignedNumber(client_nonce), + ) + .add_typed_default(DataType::Timestamp, DataValue::UnsignedNumber(new_nonce)) + .add_typed_default(DataType::Signature, DataValue::Bytes(host_sig)) + .add_typed_default(DataType::Id, DataValue::UnsignedNumber(assigned_id as u128)); + + sender.send(&response).await.ok()?; + + // 5. Version negotiation + let negotiated = self.registry.negotiate(&[client_version])?; + let codec = VersionedCodec::new(self.registry.clone()); + + Some(MTPConnection { + version: negotiated, + codec, + sender, + receiver, + auth_state: AuthState::Authenticated, + client_id: assigned_id, + client_public_key: Some(client_bundle), + }) + } +} + /* * Extract the protocol version from an initial `CommunicationValue`. * @@ -110,8 +338,11 @@ mod tests { #[test] fn version_extraction() { let tm = mtp_codec::TypeMap::latest(); - let msg = CommunicationValue::from_comm(mtp_codec::CommunicationType::Identification, &tm) - .add_data(DataTypeId(3), DataValue::Str("2.0".to_string())); + let msg = mtp_codec::CommunicationValue::from_comm( + mtp_codec::CommunicationType::Identification, + &tm, + ) + .add_data(DataTypeId(3), DataValue::Str("2.0".to_string())); let version = extract_version(&msg); assert_eq!(version, Some(Version(2, 0))); } @@ -119,15 +350,21 @@ mod tests { #[test] fn version_extraction_returns_none_for_missing() { let tm = mtp_codec::TypeMap::latest(); - let msg = CommunicationValue::from_comm(mtp_codec::CommunicationType::Identification, &tm); + let msg = mtp_codec::CommunicationValue::from_comm( + mtp_codec::CommunicationType::Identification, + &tm, + ); assert!(extract_version(&msg).is_none()); } #[test] fn version_extraction_bad_format() { let tm = mtp_codec::TypeMap::latest(); - let msg = CommunicationValue::from_comm(mtp_codec::CommunicationType::Identification, &tm) - .add_data(DataTypeId(3), DataValue::UnsignedNumber(42)); + let msg = mtp_codec::CommunicationValue::from_comm( + mtp_codec::CommunicationType::Identification, + &tm, + ) + .add_data(DataTypeId(3), DataValue::UnsignedNumber(42)); assert!(extract_version(&msg).is_none()); } } diff --git a/registry/Cargo.toml b/registry/Cargo.toml deleted file mode 100644 index 1188739..0000000 --- a/registry/Cargo.toml +++ /dev/null @@ -1,8 +0,0 @@ -[package] -name = "registry" -version = "0.1.0" -edition = "2024" - -[dependencies] -mtp-common = { path = "../common" } -mtp-type-map = { path = "../type-map", features = ["multi-version"] } diff --git a/registry/src/lib.rs b/registry/src/lib.rs deleted file mode 100644 index a738684..0000000 --- a/registry/src/lib.rs +++ /dev/null @@ -1,127 +0,0 @@ -use std::collections::BTreeMap; - -pub use mtp_type_map::{ - CommunicationType, CommunicationTypeId, DataType, DataTypeId, TypeMap, Version, -}; - -/* - * Reserved communication & data-type namespace (0..32). - * These are fixed across all protocol versions for version negotiation - * and security. They are never assigned by user configuration. - */ -pub const INTERNAL_COMM_RESERVED: std::ops::Range = 0..32; -pub const INTERNAL_DATA_RESERVED: std::ops::Range = 0..32; - -/* - * Multi-version type-map registry. - * - * Stores one `TypeMap` per protocol version and provides version - * negotiation for the host. Hosts build this from the compiled-in - * type maps (via `Registry::builtin()` or programmatically). - */ -#[derive(Clone, Debug)] -pub struct Registry { - versions: BTreeMap, -} - -impl Registry { - pub fn new() -> Self { - Self { - versions: BTreeMap::new(), - } - } - - // Register a type map, keyed by its version. - pub fn register(&mut self, typemap: TypeMap) { - self.versions.insert(typemap.version.clone(), typemap); - } - - // Look up the type map for an exact version. - pub fn get(&self, version: &Version) -> Option<&TypeMap> { - self.versions.get(version) - } - - // Check if a specific version is supported. - pub fn supports(&self, version: &Version) -> bool { - self.versions.contains_key(version) - } - - /* - * Negotiate the highest mutually supported version from a list. - * - * For single-version clients (the typical case) this is equivalent - * to checking `supports(&client_versions[0])`. - */ - pub fn negotiate(&self, client_versions: &[Version]) -> Option { - client_versions - .iter() - .filter(|v| self.versions.contains_key(v)) - .max() - .cloned() - } - - // Return the latest (highest) registered version's type map. - pub fn latest(&self) -> Option<&TypeMap> { - self.versions.last_key_value().map(|(_, v)| v) - } - - /* - * Build the registry from the compiled-in configuration. - * - * Uses the `TypeMap::vX_Y()` constructors generated by - * `type-map/build.rs` from the YAML config. - */ - pub fn builtin() -> Self { - let mut r = Self::new(); - for tm in mtp_type_map::builtin_type_maps() { - r.register(tm); - } - r - } -} - -impl Default for Registry { - fn default() -> Self { - Self::new() - } -} - -/* ================================ TESTS ================================ */ -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn builtin_contains_versions() { - let r = Registry::builtin(); - // At least the reserved-only fallback should exist - assert!(r.supports(&Version(0, 0))); - } - - #[test] - fn negotiate_finds_highest() { - let mut r = Registry::new(); - r.register(TypeMap::new(Version(0, 0))); - r.register(TypeMap::new(Version(1, 0))); - r.register(TypeMap::new(Version(2, 0))); - - let client = &[Version(0, 0)]; - assert_eq!(r.negotiate(client), Some(Version(0, 0))); - - let client = &[Version(1, 0), Version(0, 0)]; - assert_eq!(r.negotiate(client), Some(Version(1, 0))); - - let client = &[Version(5, 0)]; - assert_eq!(r.negotiate(client), None); - } - - #[test] - fn latest_returns_highest() { - let mut r = Registry::new(); - r.register(TypeMap::new(Version(0, 0))); - r.register(TypeMap::new(Version(2, 0))); - r.register(TypeMap::new(Version(1, 0))); - - assert_eq!(r.latest().unwrap().version, Version(2, 0)); - } -} diff --git a/src/lib.rs b/src/lib.rs new file mode 100644 index 0000000..7d74cea --- /dev/null +++ b/src/lib.rs @@ -0,0 +1,13 @@ +pub use mtp_common as common; +pub use mtp_type_map as type_map; +pub use mtp_codec as codec; +pub use mtp_transport as transport; + +#[cfg(feature = "crypto")] +pub use mtp_crypto as crypto; + +#[cfg(feature = "host")] +pub use mtp_host as host; + +#[cfg(feature = "client")] +pub use mtp_client as client; diff --git a/transport/Cargo.toml b/transport/Cargo.toml index 1958e5a..d7ccda3 100644 --- a/transport/Cargo.toml +++ b/transport/Cargo.toml @@ -11,14 +11,15 @@ wtransport = { version = "0.7.1", default-features = false, features = [ "quinn", "self-signed", ] } -rustls = { version = "0.23.40" } -quinn = { version = "0.11.9", default-features = false, features = [ +rustls = { version = "0.23.41" } +quinn = { version = "0.11.11", default-features = false, features = [ "rustls-aws-lc-rs", "rustls", ] } tokio = { version = "1", features = ["full"] } thiserror = "2.0.18" rustls-native-certs = "0.8.4" +log = "0.4" [features] default = [] diff --git a/transport/src/connection.rs b/transport/src/connection.rs index 3f1ac96..1738999 100644 --- a/transport/src/connection.rs +++ b/transport/src/connection.rs @@ -53,10 +53,10 @@ impl Default for Policy { } } -#[allow(unused)] enum ReceivedFrame { Message(CommunicationValue), ClosedByPeer, + #[allow(dead_code)] Idle, } @@ -215,7 +215,7 @@ impl Sender { .await .map_err(|_| CommunicationError::StreamError)? { - println!("[Sender] close frame finish failed: {e}"); + log::warn!("[Sender] close frame finish failed: {e}"); } Ok(()) @@ -445,7 +445,7 @@ impl Receiver { if e.kind() == ErrorKind::UnexpectedEof { return Ok(ReceivedFrame::Idle); } - println!("[Receiver] read_u32 failed: {e}"); + log::warn!("[Receiver] read_u32 failed: {e}"); return Err(CommunicationError::StreamError); } } @@ -468,13 +468,13 @@ impl Receiver { _ => return Err(e.into()), }, Err(_) => { - println!("[Receiver] read_exact timed out (len={})", len); + log::warn!("[Receiver] read_exact timed out (len={})", len); return Err(CommunicationError::StreamError); } } let message = CommunicationValue::from_bytes(&buf) - .ok_or(CommunicationError::ParseCommunicationValue)?; + .map_err(|_| CommunicationError::ParseCommunicationValue)?; Ok(ReceivedFrame::Message(message)) } diff --git a/transport/src/host.rs b/transport/src/host.rs index b6da6dd..d963fab 100644 --- a/transport/src/host.rs +++ b/transport/src/host.rs @@ -108,7 +108,8 @@ async fn configure_server( .ok() .and_then(|s| if s.is_empty() { None } else { Some(s) }) .unwrap_or_else(|| "::".to_string()) - .parse::()?; + .parse::() + .map_err(|e| CommunicationError::ParseError(e.to_string()))?; let bind_addr = SocketAddr::new(bind_ip, port); let server_config = ServerConfig::builder() diff --git a/type-map/Cargo.toml b/type-map/Cargo.toml index a233cc7..3dac68f 100644 --- a/type-map/Cargo.toml +++ b/type-map/Cargo.toml @@ -6,7 +6,9 @@ build = "build.rs" [features] default = [] -multi-version = [] +# Enables multi-version type-map constructors, builtin_type_maps(), and +# the Registry struct for version negotiation. Used by host, not client. +registry = [] [dependencies] diff --git a/type-map/build.rs b/type-map/build.rs index 46e7dd2..8930db0 100644 --- a/type-map/build.rs +++ b/type-map/build.rs @@ -24,45 +24,143 @@ struct ReservedEntry { } const RESERVED_COMM_TYPES: &[ReservedEntry] = &[ - ReservedEntry { name: "Error", id: 0 }, - ReservedEntry { name: "ErrorParsing", id: 1 }, - ReservedEntry { name: "ErrorBadVersion", id: 2 }, - ReservedEntry { name: "Disconnect", id: 3 }, - ReservedEntry { name: "Redirect", id: 4 }, - ReservedEntry { name: "Shutdown", id: 5 }, - ReservedEntry { name: "BadRequest", id: 6 }, - ReservedEntry { name: "Unauthorized", id: 7 }, - ReservedEntry { name: "Forbidden", id: 8 }, - ReservedEntry { name: "NotFound", id: 9 }, - ReservedEntry { name: "TooManyRequests", id: 10 }, - ReservedEntry { name: "InternalServerError", id: 11 }, - ReservedEntry { name: "BadGateway", id: 12 }, - ReservedEntry { name: "ServiceUnavailable", id: 13 }, - ReservedEntry { name: "GatewayTimeout", id: 14 }, - ReservedEntry { name: "Identification", id: 15 }, - ReservedEntry { name: "IdentificationResponse", id: 16 }, - ReservedEntry { name: "Register", id: 17 }, - ReservedEntry { name: "RegisterResponse", id: 18 }, - ReservedEntry { name: "Ping", id: 19 }, - ReservedEntry { name: "Pong", id: 20 }, + ReservedEntry { + name: "Error", + id: 0, + }, + ReservedEntry { + name: "ErrorParsing", + id: 1, + }, + ReservedEntry { + name: "ErrorBadVersion", + id: 2, + }, + ReservedEntry { + name: "Disconnect", + id: 3, + }, + ReservedEntry { + name: "Redirect", + id: 4, + }, + ReservedEntry { + name: "Shutdown", + id: 5, + }, + ReservedEntry { + name: "BadRequest", + id: 6, + }, + ReservedEntry { + name: "Unauthorized", + id: 7, + }, + ReservedEntry { + name: "Forbidden", + id: 8, + }, + ReservedEntry { + name: "NotFound", + id: 9, + }, + ReservedEntry { + name: "TooManyRequests", + id: 10, + }, + ReservedEntry { + name: "InternalServerError", + id: 11, + }, + ReservedEntry { + name: "BadGateway", + id: 12, + }, + ReservedEntry { + name: "ServiceUnavailable", + id: 13, + }, + ReservedEntry { + name: "GatewayTimeout", + id: 14, + }, + ReservedEntry { + name: "Identification", + id: 15, + }, + ReservedEntry { + name: "IdentificationResponse", + id: 16, + }, + ReservedEntry { + name: "Register", + id: 17, + }, + ReservedEntry { + name: "RegisterResponse", + id: 18, + }, + ReservedEntry { + name: "Ping", + id: 19, + }, + ReservedEntry { + name: "Pong", + id: 20, + }, ]; const RESERVED_DATA_TYPES: &[ReservedEntry] = &[ - ReservedEntry { name: "Error", id: 0 }, - ReservedEntry { name: "ErrorParsing", id: 1 }, - ReservedEntry { name: "ErrorMessage", id: 2 }, - ReservedEntry { name: "Version", id: 3 }, - ReservedEntry { name: "Description", id: 4 }, - ReservedEntry { name: "Timestamp", id: 5 }, + ReservedEntry { + name: "Error", + id: 0, + }, + ReservedEntry { + name: "ErrorParsing", + id: 1, + }, + ReservedEntry { + name: "ErrorMessage", + id: 2, + }, + ReservedEntry { + name: "Version", + id: 3, + }, + ReservedEntry { + name: "Description", + id: 4, + }, + ReservedEntry { + name: "Timestamp", + id: 5, + }, ReservedEntry { name: "Id", id: 6 }, - ReservedEntry { name: "Nonce", id: 7 }, - ReservedEntry { name: "PublicKeys", id: 8 }, - ReservedEntry { name: "Signature", id: 9 }, + ReservedEntry { + name: "ClientNonce", + id: 7, + }, + ReservedEntry { + name: "ServerNonce", + id: 8, + }, + ReservedEntry { + name: "PublicKeys", + id: 9, + }, + ReservedEntry { + name: "Signature", + id: 10, + }, + ReservedEntry { + name: "Connected", + id: 11, + }, ]; fn main() { let out = std::path::PathBuf::from(std::env::var("OUT_DIR").unwrap()); - let multi_version = std::env::var("CARGO_FEATURE_MULTI_VERSION").is_ok(); + let multi_version = std::env::var("CARGO_FEATURE_REGISTRY").is_ok(); let config = match std::env::var("MTP_TYPE_MAPS") { Ok(config_path) => { @@ -251,6 +349,33 @@ fn generate_comm_type_enum(out: &mut String, user_names: &BTreeSet<&str>) { writeln!(out, "}}").unwrap(); writeln!(out).unwrap(); + writeln!(out, "impl CommunicationType {{").unwrap(); + writeln!(out, " pub fn from_name(s: &str) -> Option {{").unwrap(); + writeln!(out, " match s {{").unwrap(); + + for entry in RESERVED_COMM_TYPES { + writeln!( + out, + " \"{}\" => Some(CommunicationType::{}),", + entry.name, entry.name + ) + .unwrap(); + } + for name in user_names { + writeln!( + out, + " \"{}\" => Some(CommunicationType::{}),", + name, name + ) + .unwrap(); + } + + writeln!(out, " _ => None,").unwrap(); + writeln!(out, " }}").unwrap(); + writeln!(out, " }}").unwrap(); + writeln!(out, "}}").unwrap(); + writeln!(out).unwrap(); + writeln!(out, "impl std::fmt::Display for CommunicationType {{").unwrap(); writeln!( out, @@ -303,6 +428,33 @@ fn generate_data_type_enum(out: &mut String, user_names: &BTreeSet<&str>) { writeln!(out, "}}").unwrap(); writeln!(out).unwrap(); + writeln!(out, "impl DataType {{").unwrap(); + writeln!(out, " pub fn from_name(s: &str) -> Option {{").unwrap(); + writeln!(out, " match s {{").unwrap(); + + for entry in RESERVED_DATA_TYPES { + writeln!( + out, + " \"{}\" => Some(DataType::{}),", + entry.name, entry.name + ) + .unwrap(); + } + for name in user_names { + writeln!( + out, + " \"{}\" => Some(DataType::{}),", + name, name + ) + .unwrap(); + } + + writeln!(out, " _ => None,").unwrap(); + writeln!(out, " }}").unwrap(); + writeln!(out, " }}").unwrap(); + writeln!(out, "}}").unwrap(); + writeln!(out).unwrap(); + writeln!(out, "impl std::fmt::Display for DataType {{").unwrap(); writeln!( out, @@ -429,11 +581,7 @@ fn generate_single_version_lookup(out: &mut String, config: &Config) { ) .unwrap(); writeln!(out, " match self.version {{").unwrap(); - writeln!( - out, - " PROTOCOL_VERSION => match ct {{" - ) - .unwrap(); + writeln!(out, " PROTOCOL_VERSION => match ct {{").unwrap(); for entry in RESERVED_COMM_TYPES { writeln!( out, @@ -466,11 +614,7 @@ fn generate_single_version_lookup(out: &mut String, config: &Config) { ) .unwrap(); writeln!(out, " match self.version {{").unwrap(); - writeln!( - out, - " PROTOCOL_VERSION => match dt {{" - ) - .unwrap(); + writeln!(out, " PROTOCOL_VERSION => match dt {{").unwrap(); for entry in RESERVED_DATA_TYPES { writeln!( out, @@ -529,7 +673,11 @@ fn generate_reverse_lookups( } } - writeln!(out, "pub fn communication_type_name(id: u16) -> Option<&'static str> {{").unwrap(); + writeln!( + out, + "pub fn communication_type_name(id: u16) -> Option<&'static str> {{" + ) + .unwrap(); writeln!(out, " match id {{").unwrap(); for (id, name) in &id_to_comm { writeln!(out, " {} => Some(\"{}\"),", id, name).unwrap(); @@ -557,7 +705,11 @@ fn generate_reverse_lookups( } } - writeln!(out, "pub fn data_type_name(id: u16) -> Option<&'static str> {{").unwrap(); + writeln!( + out, + "pub fn data_type_name(id: u16) -> Option<&'static str> {{" + ) + .unwrap(); writeln!(out, " match id {{").unwrap(); for (id, name) in &id_to_data { writeln!(out, " {} => Some(\"{}\"),", id, name).unwrap(); @@ -570,7 +722,11 @@ fn generate_reverse_lookups( fn generate_enum_conversion_methods(out: &mut String) { writeln!(out, "impl CommunicationType {{").unwrap(); - writeln!(out, " pub fn to_id(self, tm: &TypeMap) -> CommunicationTypeId {{").unwrap(); + writeln!( + out, + " pub fn to_id(self, tm: &TypeMap) -> CommunicationTypeId {{" + ) + .unwrap(); writeln!( out, " CommunicationTypeId(tm.comm_id_enum(self).expect(\"CommunicationType not found in this TypeMap version\"))" diff --git a/type-map/src/lib.rs b/type-map/src/lib.rs index 3173de2..29f5429 100644 --- a/type-map/src/lib.rs +++ b/type-map/src/lib.rs @@ -1,9 +1,40 @@ +pub const INTERNAL_COMM_RESERVED: std::ops::Range = 0..32; +pub const INTERNAL_DATA_RESERVED: std::ops::Range = 0..32; + #[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)] pub struct CommunicationTypeId(pub u16); +impl CommunicationTypeId { + pub fn try_new(id: u16) -> Result { + if INTERNAL_COMM_RESERVED.contains(&id) { + Err(()) + } else { + Ok(Self(id)) + } + } + + pub fn is_reserved(&self) -> bool { + INTERNAL_COMM_RESERVED.contains(&self.0) + } +} + #[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)] pub struct DataTypeId(pub u16); +impl DataTypeId { + pub fn try_new(id: u16) -> Result { + if INTERNAL_DATA_RESERVED.contains(&id) { + Err(()) + } else { + Ok(Self(id)) + } + } + + pub fn is_reserved(&self) -> bool { + INTERNAL_DATA_RESERVED.contains(&self.0) + } +} + #[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] pub struct Version(pub u16, pub u16); @@ -43,3 +74,109 @@ impl TypeMap { } include!(concat!(env!("OUT_DIR"), "/types.rs")); + +/* ============================= REGISTRY ============================= */ +#[cfg(feature = "registry")] +pub use registry::*; + +#[cfg(feature = "registry")] +mod registry { + use std::collections::BTreeMap; + use crate::{Version, TypeMap, builtin_type_maps}; + + /* + * Multi-version type-map registry. + * + * Stores one `TypeMap` per protocol version and provides version + * negotiation for the host. Hosts build this from the compiled-in + * type maps (via `Registry::builtin()` or programmatically). + */ + #[derive(Clone, Debug)] + pub struct Registry { + versions: BTreeMap, + } + + impl Registry { + pub fn new() -> Self { + Self { + versions: BTreeMap::new(), + } + } + + pub fn register(&mut self, typemap: TypeMap) { + self.versions.insert(typemap.version.clone(), typemap); + } + + pub fn get(&self, version: &Version) -> Option<&TypeMap> { + self.versions.get(version) + } + + pub fn supports(&self, version: &Version) -> bool { + self.versions.contains_key(version) + } + + pub fn negotiate(&self, client_versions: &[Version]) -> Option { + client_versions + .iter() + .filter(|v| self.versions.contains_key(v)) + .max() + .cloned() + } + + pub fn latest(&self) -> Option<&TypeMap> { + self.versions.last_key_value().map(|(_, v)| v) + } + + pub fn builtin() -> Self { + let mut r = Self::new(); + for tm in builtin_type_maps() { + r.register(tm); + } + r + } + } + + impl Default for Registry { + fn default() -> Self { + Self::new() + } + } + + #[cfg(test)] + mod tests { + use super::*; + + #[test] + fn builtin_contains_versions() { + let r = Registry::builtin(); + assert!(r.supports(&Version(0, 0))); + } + + #[test] + fn negotiate_finds_highest() { + let mut r = Registry::new(); + r.register(TypeMap::new(Version(0, 0))); + r.register(TypeMap::new(Version(1, 0))); + r.register(TypeMap::new(Version(2, 0))); + + let client = &[Version(0, 0)]; + assert_eq!(r.negotiate(client), Some(Version(0, 0))); + + let client = &[Version(1, 0), Version(0, 0)]; + assert_eq!(r.negotiate(client), Some(Version(1, 0))); + + let client = &[Version(5, 0)]; + assert_eq!(r.negotiate(client), None); + } + + #[test] + fn latest_returns_highest() { + let mut r = Registry::new(); + r.register(TypeMap::new(Version(0, 0))); + r.register(TypeMap::new(Version(2, 0))); + r.register(TypeMap::new(Version(1, 0))); + + assert_eq!(r.latest().unwrap().version, Version(2, 0)); + } + } +}