use parking_lot::RwLock;
use serde_json::Value;
use std::collections::HashMap;
use std::sync::Arc;
pub trait SessionStore: Send + Sync {
fn read(&self, session_id: &str) -> Option<HashMap<String, Value>>;
fn write(&self, session_id: &str, data: HashMap<String, Value>);
fn destroy(&self, session_id: &str);
fn exists(&self, session_id: &str) -> bool {
self.read(session_id).is_some()
}
}
#[derive(Debug, Clone, Default)]
pub struct MemorySessionStore {
data: Arc<RwLock<HashMap<String, HashMap<String, Value>>>>,
}
impl MemorySessionStore {
pub fn new() -> Self {
Self::default()
}
}
impl SessionStore for MemorySessionStore {
fn read(&self, session_id: &str) -> Option<HashMap<String, Value>> {
self.data.read().get(session_id).cloned()
}
fn write(&self, session_id: &str, data: HashMap<String, Value>) {
self.data.write().insert(session_id.to_string(), data);
}
fn destroy(&self, session_id: &str) {
self.data.write().remove(session_id);
}
fn exists(&self, session_id: &str) -> bool {
self.data.read().contains_key(session_id)
}
}
const FLASH_PREFIX: &str = "__flash__:";
pub struct Session {
session_id: String,
prefix: String,
store: Arc<dyn SessionStore>,
}
impl Session {
pub fn new(session_id: impl Into<String>, store: impl SessionStore + 'static) -> Self {
Self {
session_id: session_id.into(),
prefix: String::new(),
store: Arc::new(store),
}
}
pub fn with_shared_store(session_id: impl Into<String>, store: Arc<dyn SessionStore>) -> Self {
Self {
session_id: session_id.into(),
prefix: String::new(),
store,
}
}
#[must_use]
pub fn with_prefix(mut self, prefix: impl Into<String>) -> Self {
self.prefix = prefix.into();
self
}
pub fn session_id(&self) -> &str {
&self.session_id
}
fn full_key(&self, name: &str) -> String {
if self.prefix.is_empty() {
name.to_string()
} else {
format!("{}{}", self.prefix, name)
}
}
pub fn set(&self, name: &str, value: Value) {
let mut data = self.store.read(&self.session_id).unwrap_or_default();
data.insert(self.full_key(name), value);
self.store.write(&self.session_id, data);
}
pub fn get(&self, name: &str) -> Option<Value> {
let data = self.store.read(&self.session_id)?;
data.get(&self.full_key(name)).cloned()
}
pub fn get_with_default(&self, name: &str, default: Value) -> Value {
self.get(name).unwrap_or(default)
}
pub fn has(&self, name: &str) -> bool {
self.get(name).is_some()
}
pub fn delete(&self, name: &str) -> Option<Value> {
let mut data = self.store.read(&self.session_id)?;
let key = self.full_key(name);
let removed = data.remove(&key);
self.store.write(&self.session_id, data);
removed
}
pub fn clear(&self) {
self.store.destroy(&self.session_id);
}
pub fn flash(&self, name: &str, value: Value) {
let flash_key = format!("{}{}", FLASH_PREFIX, name);
self.set(&flash_key, value);
}
pub fn get_flash(&self, name: &str) -> Option<Value> {
let flash_key = format!("{}{}", FLASH_PREFIX, name);
self.get(&flash_key)
}
pub fn clear_flash(&self) {
let mut data = match self.store.read(&self.session_id) {
Some(d) => d,
None => return,
};
let flash_keys: Vec<String> = data
.keys()
.filter(|k| k.starts_with(FLASH_PREFIX))
.cloned()
.collect();
for key in flash_keys {
data.remove(&key);
}
self.store.write(&self.session_id, data);
}
pub fn flush(&self) {
self.clear();
}
pub fn all(&self) -> HashMap<String, Value> {
let data = self.store.read(&self.session_id).unwrap_or_default();
data.into_iter()
.filter(|(k, _)| !k.starts_with(FLASH_PREFIX))
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_memory_store_write_read_roundtrip() {
let store = MemorySessionStore::new();
let mut data = HashMap::new();
data.insert("user_id".to_string(), json!(12345));
data.insert("name".to_string(), json!("alice"));
store.write("session-1", data.clone());
let read = store.read("session-1").unwrap();
assert_eq!(read.len(), 2);
assert_eq!(read.get("user_id"), Some(&json!(12345)));
assert_eq!(read.get("name"), Some(&json!("alice")));
}
#[test]
fn test_memory_store_read_nonexistent_returns_none() {
let store = MemorySessionStore::new();
assert!(store.read("nonexistent").is_none());
}
#[test]
fn test_memory_store_destroy() {
let store = MemorySessionStore::new();
let data = HashMap::new();
store.write("session-1", data);
assert!(store.exists("session-1"));
store.destroy("session-1");
assert!(!store.exists("session-1"));
}
#[test]
fn test_memory_store_isolated_by_session_id() {
let store = MemorySessionStore::new();
let mut data1 = HashMap::new();
data1.insert("user".to_string(), json!("alice"));
store.write("session-1", data1);
let mut data2 = HashMap::new();
data2.insert("user".to_string(), json!("bob"));
store.write("session-2", data2);
assert_eq!(
store.read("session-1").unwrap().get("user"),
Some(&json!("alice"))
);
assert_eq!(
store.read("session-2").unwrap().get("user"),
Some(&json!("bob"))
);
}
#[test]
fn test_memory_store_overwrite() {
let store = MemorySessionStore::new();
let mut data = HashMap::new();
data.insert("key".to_string(), json!("old"));
store.write("session-1", data);
let mut new_data = HashMap::new();
new_data.insert("key".to_string(), json!("new"));
store.write("session-1", new_data);
assert_eq!(
store.read("session-1").unwrap().get("key"),
Some(&json!("new"))
);
}
fn make_session() -> Session {
Session::new("test-session-id", MemorySessionStore::new())
}
#[test]
fn test_session_set_get() {
let session = make_session();
session.set("user_id", json!(12345));
assert_eq!(session.get("user_id"), Some(json!(12345)));
}
#[test]
fn test_session_set_string_value() {
let session = make_session();
session.set("name", json!("alice"));
assert_eq!(session.get("name"), Some(json!("alice")));
}
#[test]
fn test_session_set_object_value() {
let session = make_session();
session.set("user", json!({"id": 1, "name": "bob"}));
let value = session.get("user").unwrap();
assert_eq!(value["id"], 1);
assert_eq!(value["name"], "bob");
}
#[test]
fn test_session_get_nonexistent_returns_none() {
let session = make_session();
assert_eq!(session.get("missing"), None);
}
#[test]
fn test_session_get_with_default_returns_value_when_exists() {
let session = make_session();
session.set("key", json!("actual"));
assert_eq!(
session.get_with_default("key", json!("default")),
json!("actual")
);
}
#[test]
fn test_session_get_with_default_returns_default_when_missing() {
let session = make_session();
assert_eq!(
session.get_with_default("missing", json!("default")),
json!("default")
);
}
#[test]
fn test_session_has_existing_key() {
let session = make_session();
session.set("key", json!(1));
assert!(session.has("key"));
}
#[test]
fn test_session_has_nonexistent_key() {
let session = make_session();
assert!(!session.has("missing"));
}
#[test]
fn test_session_delete_returns_value() {
let session = make_session();
session.set("key", json!("value"));
let removed = session.delete("key");
assert_eq!(removed, Some(json!("value")));
assert!(!session.has("key"));
}
#[test]
fn test_session_delete_nonexistent_returns_none() {
let session = make_session();
let removed = session.delete("missing");
assert_eq!(removed, None);
}
#[test]
fn test_session_clear_removes_all_data() {
let session = make_session();
session.set("key1", json!(1));
session.set("key2", json!(2));
session.set("key3", json!(3));
session.clear();
assert!(!session.has("key1"));
assert!(!session.has("key2"));
assert!(!session.has("key3"));
}
#[test]
fn test_session_all_returns_non_flash_data() {
let session = make_session();
session.set("key1", json!(1));
session.set("key2", json!("two"));
session.flash("temp", json!("flash"));
let all = session.all();
assert_eq!(all.len(), 2); assert_eq!(all.get("key1"), Some(&json!(1)));
assert_eq!(all.get("key2"), Some(&json!("two")));
}
#[test]
fn test_session_prefix_isolation() {
let store = MemorySessionStore::new();
let session_a = Session::new("sid", store.clone()).with_prefix("app_a_");
let session_b = Session::new("sid", store.clone()).with_prefix("app_b_");
session_a.set("user", json!("alice"));
session_b.set("user", json!("bob"));
assert_eq!(session_a.get("user"), Some(json!("alice")));
assert_eq!(session_b.get("user"), Some(json!("bob")));
}
#[test]
fn test_session_prefix_empty_by_default() {
let session = make_session();
assert_eq!(session.prefix, "");
}
#[test]
fn test_session_flash_set_get() {
let session = make_session();
session.flash("success", json!("操作成功"));
assert_eq!(session.get_flash("success"), Some(json!("操作成功")));
}
#[test]
fn test_session_flash_not_in_regular_get() {
let session = make_session();
session.flash("temp", json!("flash data"));
assert_eq!(session.get("temp"), None);
assert_eq!(session.get("__flash__:temp"), Some(json!("flash data")));
}
#[test]
fn test_session_clear_flash_removes_flash_data() {
let session = make_session();
session.flash("temp1", json!(1));
session.flash("temp2", json!(2));
session.set("regular", json!("keep"));
session.clear_flash();
assert_eq!(session.get_flash("temp1"), None);
assert_eq!(session.get_flash("temp2"), None);
assert_eq!(session.get("regular"), Some(json!("keep")));
}
#[test]
fn test_session_clear_flash_when_no_data() {
let session = make_session();
session.clear_flash();
}
#[test]
fn test_session_flush_equals_clear() {
let session1 = make_session();
let session2 = make_session();
session1.set("key", json!(1));
session2.set("key", json!(1));
session1.clear();
session2.flush();
assert!(!session1.has("key"));
assert!(!session2.has("key"));
}
#[test]
fn test_multiple_sessions_share_store() {
let store = Arc::new(MemorySessionStore::new());
let session1 = Session::with_shared_store("sid-1", store.clone());
let session2 = Session::with_shared_store("sid-2", store.clone());
session1.set("user", json!("alice"));
session2.set("user", json!("bob"));
assert_eq!(session1.get("user"), Some(json!("alice")));
assert_eq!(session2.get("user"), Some(json!("bob")));
session1.clear();
assert!(!session1.has("user"));
assert!(session2.has("user"));
}
#[test]
fn test_session_id_access() {
let session = Session::new("my-session-id", MemorySessionStore::new());
assert_eq!(session.session_id(), "my-session-id");
}
#[test]
fn test_php_consistency_session_full_flow() {
let store = Arc::new(MemorySessionStore::new());
let login_session = Session::with_shared_store("sid-login", store.clone());
login_session.set(
"szshop_clerk",
json!({"clerk_id": 100, "name": "张三", "store_id": 5}),
);
let later_session = Session::with_shared_store("sid-login", store.clone());
let clerk = later_session.get("szshop_clerk").unwrap();
assert_eq!(clerk["clerk_id"], 100);
assert_eq!(clerk["name"], "张三");
assert_eq!(clerk["store_id"], 5);
later_session.clear();
assert!(!later_session.has("szshop_clerk"));
}
#[test]
fn test_php_consistency_flash_message_flow() {
let store = Arc::new(MemorySessionStore::new());
let submit_session = Session::with_shared_store("sid", store.clone());
submit_session.flash("success", json!("保存成功"));
let redirect_session = Session::with_shared_store("sid", store.clone());
assert_eq!(
redirect_session.get_flash("success"),
Some(json!("保存成功"))
);
redirect_session.clear_flash();
let next_session = Session::with_shared_store("sid", store.clone());
assert_eq!(next_session.get_flash("success"), None);
}
}