auto-register messages via inventory

message_registry! is gone, the registry is a process-wide global now.
This commit is contained in:
scroll 2026-08-23 07:56:09 +03:00
parent 323450a977
commit 2151d4ccbe
10 changed files with 74 additions and 44 deletions

16
Cargo.lock generated
View file

@ -103,6 +103,15 @@ dependencies = [
"wasi", "wasi",
] ]
[[package]]
name = "inventory"
version = "0.3.24"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a4f0c30c76f2f4ccee3fe55a2435f691ca00c0e4bd87abe4f4a851b1d4dac39b"
dependencies = [
"rustversion",
]
[[package]] [[package]]
name = "itoa" name = "itoa"
version = "1.0.18" version = "1.0.18"
@ -257,6 +266,12 @@ version = "0.8.11"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4"
[[package]]
name = "rustversion"
version = "1.0.23"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f"
[[package]] [[package]]
name = "scroll-server" name = "scroll-server"
version = "0.1.0" version = "0.1.0"
@ -419,6 +434,7 @@ dependencies = [
name = "titan" name = "titan"
version = "0.1.0" version = "0.1.0"
dependencies = [ dependencies = [
"inventory",
"thiserror", "thiserror",
"titan-derive", "titan-derive",
"tokio", "tokio",

View file

@ -34,6 +34,7 @@ serde = { version = "1", features = ["derive"] }
serde_json = "1" serde_json = "1"
rand = "0.8" rand = "0.8"
async-trait = "0.1" async-trait = "0.1"
inventory = "0.3"
proc-macro2 = "1" proc-macro2 = "1"
quote = "1" quote = "1"
syn = { version = "2", features = ["full", "extra-traits"] } syn = { version = "2", features = ["full", "extra-traits"] }

View file

@ -1,36 +1,11 @@
use titan::{message_registry, MessageRegistry}; use std::sync::Arc;
use crate::messages::{ use titan::MessageRegistry;
AskForAvatarStreamMessage, AskForBattleReplayStreamMessage,
AskForPlayingFacebookFriendsMessage, AskForTVContentMessage, AvailableServerCommandMessage,
ClientCapabilitiesMessage, EndClientTurnMessage, GoHomeMessage, KeepAliveMessage,
KeepAliveServerMessage, LoginFailedMessage, LoginMessage, LoginOkMessage, OwnHomeDataMessage,
OutOfSyncMessage, ServerErrorMessage, SetDeviceTokenMessage, StartMissionMessage,
};
pub struct LogicScrollMessageFactory; pub struct LogicScrollMessageFactory;
impl LogicScrollMessageFactory { impl LogicScrollMessageFactory {
pub fn registry() -> MessageRegistry { pub fn registry() -> Arc<MessageRegistry> {
scroll_message_registry() scroll_message_registry()
} }
} }
pub fn scroll_message_registry() -> MessageRegistry { pub fn scroll_message_registry() -> Arc<MessageRegistry> {
message_registry![ MessageRegistry::global()
LoginMessage,
ClientCapabilitiesMessage,
KeepAliveMessage,
GoHomeMessage,
EndClientTurnMessage,
StartMissionMessage,
SetDeviceTokenMessage,
AskForPlayingFacebookFriendsMessage,
AskForTVContentMessage,
AskForAvatarStreamMessage,
AskForBattleReplayStreamMessage,
LoginOkMessage,
LoginFailedMessage,
KeepAliveServerMessage,
OwnHomeDataMessage,
AvailableServerCommandMessage,
OutOfSyncMessage,
ServerErrorMessage,
]
} }

View file

@ -66,6 +66,10 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
} }
Err(error) => println!("no data tables ({error}), printing raw ids"), Err(error) => println!("no data tables ({error}), printing raw ids"),
} }
println!(
"message registry: {} type(s)",
logic::scroll_message_registry().len()
);
println!("connecting to {endpoint}"); println!("connecting to {endpoint}");
let mut probe = Probe::connect(&endpoint).await?; let mut probe = Probe::connect(&endpoint).await?;
let login = LoginMessage { let login = LoginMessage {

View file

@ -23,7 +23,7 @@ impl Session {
config: Arc<GatewayConfig>, config: Arc<GatewayConfig>,
backends: Arc<Backends>, backends: Arc<Backends>,
) -> Result<(), SessionError> { ) -> Result<(), SessionError> {
let registry: Arc<MessageRegistry> = Arc::new(logic::scroll_message_registry()); let registry: Arc<MessageRegistry> = logic::scroll_message_registry();
let messaging_config = MessagingConfig { let messaging_config = MessagingConfig {
read_timeout: config.read_timeout, read_timeout: config.read_timeout,
max_payload_len: config.max_payload_len, max_payload_len: config.max_payload_len,

View file

@ -20,9 +20,11 @@ pub fn derive_message(input: TokenStream) -> TokenStream {
Ok(tokens) => tokens, Ok(tokens) => tokens,
Err(error) => return error.into_compile_error().into(), Err(error) => return error.into_compile_error().into(),
}; };
let registration = meta::expand_registration(&input);
quote::quote! { quote::quote! {
#payload #payload
#meta #meta
#registration
} }
.into() .into()
} }

View file

@ -87,3 +87,14 @@ fn parse_message_attr(input: &DeriveInput) -> Result<MessageAttr> {
name, name,
}) })
} }
pub fn expand_registration(input: &DeriveInput) -> TokenStream {
if !input.generics.params.is_empty() {
return TokenStream::new();
}
let ident = &input.ident;
quote! {
::titan::inventory::submit! {
::titan::factory::RegistryEntry::of::<#ident>()
}
}
}

View file

@ -9,5 +9,6 @@ description = "Engine layer of the scroll server: byte streams, checksum encoder
[dependencies] [dependencies]
titan-derive = { workspace = true } titan-derive = { workspace = true }
thiserror = { workspace = true } thiserror = { workspace = true }
inventory = { workspace = true }
tokio = { workspace = true } tokio = { workspace = true }
tracing = { workspace = true } tracing = { workspace = true }

View file

@ -1,7 +1,8 @@
use std::collections::HashMap; use std::collections::HashMap;
use std::fmt::Debug; use std::fmt::Debug;
use std::sync::{Arc, OnceLock};
use crate::error::{Error, Result}; use crate::error::{Error, Result};
use crate::message::{Direction, Message, MessageMeta}; use crate::message::{Direction, Message, MessageMeta, Payload};
type DecodeFn = fn(&[u8]) -> Result<Box<dyn Message>>; type DecodeFn = fn(&[u8]) -> Result<Box<dyn Message>>;
#[derive(Clone, Copy)] #[derive(Clone, Copy)]
pub struct RegistryEntry { pub struct RegistryEntry {
@ -20,7 +21,7 @@ impl Debug for RegistryEntry {
} }
} }
impl RegistryEntry { impl RegistryEntry {
pub fn of<T>() -> Self pub const fn of<T>() -> Self
where where
T: MessageMeta + Debug + Send + Sync + 'static, T: MessageMeta + Debug + Send + Sync + 'static,
{ {
@ -28,13 +29,21 @@ impl RegistryEntry {
message_type: T::MESSAGE_TYPE, message_type: T::MESSAGE_TYPE,
name: T::NAME, name: T::NAME,
direction: T::DIRECTION, direction: T::DIRECTION,
decode: |bytes| Ok(Box::new(T::from_bytes(bytes)?)), decode: decode_into_box::<T>,
} }
} }
pub fn decode(&self, payload: &[u8]) -> Result<Box<dyn Message>> { pub fn decode(&self, payload: &[u8]) -> Result<Box<dyn Message>> {
(self.decode)(payload) (self.decode)(payload)
} }
} }
fn decode_into_box<T>(payload: &[u8]) -> Result<Box<dyn Message>>
where
T: MessageMeta + Debug + Send + Sync + 'static,
{
Ok(Box::new(T::from_bytes(payload)?))
}
inventory::collect!(RegistryEntry);
static GLOBAL: OnceLock<Arc<MessageRegistry>> = OnceLock::new();
pub trait MessageFactory: Send + Sync { pub trait MessageFactory: Send + Sync {
fn create(&self, message_type: u16, message_version: u16, payload: &[u8]) -> Result<Box<dyn Message>>; fn create(&self, message_type: u16, message_version: u16, payload: &[u8]) -> Result<Box<dyn Message>>;
fn lookup(&self, message_type: u16) -> Option<&RegistryEntry>; fn lookup(&self, message_type: u16) -> Option<&RegistryEntry>;
@ -54,8 +63,25 @@ impl MessageRegistry {
} }
registry registry
} }
pub fn collect() -> Self {
Self::from_entries(inventory::iter::<RegistryEntry>.into_iter().copied())
}
pub fn global() -> Arc<MessageRegistry> {
Arc::clone(GLOBAL.get_or_init(|| {
let registry = Self::collect();
tracing::debug!(messages = registry.len(), "message registry collected");
Arc::new(registry)
}))
}
pub fn insert(&mut self, entry: RegistryEntry) { pub fn insert(&mut self, entry: RegistryEntry) {
self.entries.insert(entry.message_type, entry); if let Some(previous) = self.entries.insert(entry.message_type, entry) {
tracing::warn!(
message_type = entry.message_type,
previous = previous.name,
replacement = entry.name,
"two messages declare the same type"
);
}
} }
pub fn len(&self) -> usize { pub fn len(&self) -> usize {
self.entries.len() self.entries.len()
@ -76,11 +102,3 @@ impl MessageFactory for MessageRegistry {
self.entries.get(&message_type) self.entries.get(&message_type)
} }
} }
#[macro_export]
macro_rules! message_registry {
($($ty:ty),* $(,)?) => {
$crate::factory::MessageRegistry::from_entries([
$($crate::factory::RegistryEntry::of::<$ty>()),*
])
};
}

View file

@ -26,6 +26,7 @@ pub use json::JsonValue;
pub use logic_long::LogicLong; pub use logic_long::LogicLong;
pub use message::{Direction, Message, MessageMeta, Payload}; pub use message::{Direction, Message, MessageMeta, Payload};
pub use net::{Incoming, Messaging, MessagingConfig, MessagingError, MessagingSender, OutboundMessage}; pub use net::{Incoming, Messaging, MessagingConfig, MessagingError, MessagingSender, OutboundMessage};
pub use inventory;
pub use titan_derive::{Message, Payload}; pub use titan_derive::{Message, Payload};
pub mod prelude { pub mod prelude {
pub use crate::codec::{ pub use crate::codec::{
@ -37,5 +38,6 @@ pub mod prelude {
pub use crate::io::{ByteStreamReader, ByteStreamWriter}; pub use crate::io::{ByteStreamReader, ByteStreamWriter};
pub use crate::logic_long::LogicLong; pub use crate::logic_long::LogicLong;
pub use crate::message::{Direction, Message, MessageMeta, Payload}; pub use crate::message::{Direction, Message, MessageMeta, Payload};
pub use titan_derive::{Message, Payload}; pub use inventory;
pub use titan_derive::{Message, Payload};
} }