use std::cell::Cell;
use std::fmt::Display;
use std::sync::Arc;
use freenet_stdlib::{
client_api::{ClientError, ClientRequest, HostResponse},
prelude::ContractInstanceId,
};
use serde::{Deserialize, Serialize};
use tokio::sync::mpsc;
use crate::config::GlobalRng;
use crate::wasm_runtime::UserSecretContext;
pub type HostResult = Result<HostResponse, ClientError>;
#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
#[repr(transparent)]
pub struct RequestId(u64);
const COUNTER_BLOCK: u64 = 1_000_000;
thread_local! {
static REQUEST_ID_COUNTER: Cell<u64> = {
let idx = crate::config::GlobalRng::thread_index();
Cell::new(1 + idx * COUNTER_BLOCK)
};
}
impl RequestId {
pub fn new() -> Self {
Self(REQUEST_ID_COUNTER.with(|c| {
let v = c.get();
c.set(v + 1);
v
}))
}
pub fn reset_counter() {
let idx = crate::config::GlobalRng::thread_index();
REQUEST_ID_COUNTER.with(|c| c.set(1 + idx * COUNTER_BLOCK));
}
}
impl Default for RequestId {
fn default() -> Self {
Self::new()
}
}
impl std::fmt::Display for RequestId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "req-{}", self.0)
}
}
#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
#[repr(transparent)]
pub struct ClientId(pub(crate) usize);
impl From<ClientId> for usize {
fn from(val: ClientId) -> Self {
val.0
}
}
thread_local! {
static CLIENT_ID_COUNTER: Cell<usize> = {
let idx = crate::config::GlobalRng::thread_index();
Cell::new(1 + (idx as usize) * (COUNTER_BLOCK as usize))
};
}
impl ClientId {
pub const FIRST: Self = ClientId(0);
pub fn next() -> Self {
ClientId(CLIENT_ID_COUNTER.with(|c| {
let v = c.get();
c.set(v + 1);
v
}))
}
pub fn reset_counter() {
let idx = crate::config::GlobalRng::thread_index();
CLIENT_ID_COUNTER.with(|c| c.set(1 + (idx as usize) * (COUNTER_BLOCK as usize)));
}
}
impl Display for ClientId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.0)
}
}
pub(crate) type HostIncomingMsg = Result<OpenRequest<'static>, ClientError>;
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct AuthToken(#[serde(deserialize_with = "AuthToken::deser_auth_token")] Arc<str>);
impl AuthToken {
pub fn as_str(&self) -> &str {
&self.0
}
pub fn generate() -> AuthToken {
let mut token = [0u8; 32];
GlobalRng::fill_bytes(&mut token);
let token_str = bs58::encode(token).into_string();
AuthToken::from(token_str)
}
}
impl std::ops::Deref for AuthToken {
type Target = str;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl AuthToken {
fn deser_auth_token<'de, D>(deser: D) -> Result<Arc<str>, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = <String as Deserialize>::deserialize(deser)?;
Ok(value.into())
}
}
impl From<String> for AuthToken {
fn from(value: String) -> Self {
Self(value.into())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ConnectionScope {
Local,
#[default]
Remote,
}
impl ConnectionScope {
pub fn is_local(self) -> bool {
matches!(self, Self::Local)
}
pub fn from_source_ip(source_ip: Option<std::net::IpAddr>) -> Self {
match source_ip {
Some(ip) if is_loopback_source(ip) => Self::Local,
_ => Self::Remote,
}
}
}
pub(crate) fn is_loopback_source(ip: std::net::IpAddr) -> bool {
match ip {
std::net::IpAddr::V4(v4) => v4.is_loopback(),
std::net::IpAddr::V6(v6) => match v6.to_ipv4_mapped() {
Some(v4) => v4.is_loopback(),
None => v6.is_loopback(),
},
}
}
#[non_exhaustive]
pub struct OpenRequest<'a> {
pub client_id: ClientId,
pub request_id: RequestId,
pub request: Box<ClientRequest<'a>>,
pub notification_channel: Option<mpsc::Sender<HostResult>>,
pub token: Option<AuthToken>,
pub origin_contract: Option<ContractInstanceId>,
pub connection_scope: ConnectionScope,
pub user_context: Option<UserSecretContext>,
}
impl Display for OpenRequest<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"client request {{ client: {}, request_id: {}, req: {} }}",
&self.client_id, &self.request_id, &*self.request
)
}
}
impl<'a> OpenRequest<'a> {
pub fn into_owned(self) -> OpenRequest<'static> {
OpenRequest {
request: Box::new(self.request.into_owned()),
..self
}
}
pub fn new(id: ClientId, request: Box<ClientRequest<'a>>) -> Self {
Self {
client_id: id,
request_id: RequestId::new(),
request,
notification_channel: None,
token: None,
origin_contract: None,
connection_scope: ConnectionScope::default(),
user_context: None,
}
}
pub fn with_notification(mut self, ch: mpsc::Sender<HostResult>) -> Self {
self.notification_channel = Some(ch);
self
}
pub fn with_token(mut self, token: Option<AuthToken>) -> Self {
self.token = token;
self
}
pub fn with_origin_contract(mut self, contract: Option<ContractInstanceId>) -> Self {
self.origin_contract = contract;
self
}
pub fn with_connection_scope(mut self, scope: ConnectionScope) -> Self {
self.connection_scope = scope;
self
}
pub fn with_user_context(mut self, user_context: Option<UserSecretContext>) -> Self {
self.user_context = user_context;
self
}
}
#[cfg(test)]
mod connection_scope_tests {
use super::*;
use std::net::IpAddr;
#[test]
fn missing_source_ip_is_remote() {
assert_eq!(
ConnectionScope::from_source_ip(None),
ConnectionScope::Remote
);
assert!(!ConnectionScope::from_source_ip(None).is_local());
}
#[test]
fn loopback_sources_are_local_including_ipv4_mapped() {
for ip in [
"127.0.0.1",
"127.4.5.6",
"::1",
"::ffff:127.0.0.1",
] {
let parsed: IpAddr = ip.parse().expect("test address must parse");
assert_eq!(
ConnectionScope::from_source_ip(Some(parsed)),
ConnectionScope::Local,
"{ip} must classify as Local"
);
}
}
#[test]
fn off_host_sources_are_remote() {
for ip in [
"192.168.1.50",
"10.0.0.7",
"5.9.111.215",
"2001:db8::1",
"::ffff:192.168.1.50",
] {
let parsed: IpAddr = ip.parse().expect("test address must parse");
assert_eq!(
ConnectionScope::from_source_ip(Some(parsed)),
ConnectionScope::Remote,
"{ip} must classify as Remote"
);
}
}
#[test]
fn default_scope_is_remote() {
assert_eq!(ConnectionScope::default(), ConnectionScope::Remote);
assert!(
!OpenRequest::new(
ClientId::FIRST,
Box::new(freenet_stdlib::client_api::ClientRequest::Close),
)
.connection_scope
.is_local()
);
}
}