1use crate::base::{memory_variables_to_messages, BaseMemory};
37use crate::buffer::ConversationBufferMemory;
38use async_trait::async_trait;
39use lc_core::language_models::{BaseChatModel, LLMResult};
40use lc_core::runnables::{LcelError, Runnable, RunnableConfig};
41use lc_schema::Message;
42use std::collections::{HashMap, VecDeque};
43use std::sync::Arc;
44use tokio::sync::Mutex;
45
46pub type SharedMemory = Arc<Mutex<Box<dyn BaseMemory>>>;
48
49const DEFAULT_MAX_SESSIONS: usize = 100;
51
52pub struct RunnableWithMessageHistory<L> {
54 llm: Arc<L>,
55 mode: HistoryMode,
56}
57
58enum HistoryMode {
60 Shared(SharedMemory),
62 Sessions {
64 factory: Arc<dyn Fn(&str) -> SharedMemory + Send + Sync>,
65 cache: Mutex<SessionCache>,
67 max_sessions: usize,
69 default: SharedMemory,
71 },
72}
73
74struct SessionCache {
76 slots: HashMap<String, SharedMemory>,
77 order: VecDeque<String>,
78}
79
80impl<L> RunnableWithMessageHistory<L> {
81 pub fn new(llm: L, memory: impl BaseMemory + 'static) -> Self {
83 Self {
84 llm: Arc::new(llm),
85 mode: HistoryMode::Shared(Arc::new(Mutex::new(Box::new(memory)))),
86 }
87 }
88
89 pub fn with_session_history<F>(llm: L, factory: F) -> Self
107 where
108 F: Fn(&str) -> SharedMemory + Send + Sync + 'static,
109 {
110 Self {
111 llm: Arc::new(llm),
112 mode: HistoryMode::Sessions {
113 factory: Arc::new(factory),
114 cache: Mutex::new(SessionCache {
115 slots: HashMap::new(),
116 order: VecDeque::new(),
117 }),
118 max_sessions: DEFAULT_MAX_SESSIONS,
119 default: Arc::new(Mutex::new(Box::new(ConversationBufferMemory::new()))),
120 },
121 }
122 }
123
124 pub fn with_max_sessions(mut self, max: usize) -> Self {
127 if let HistoryMode::Sessions { max_sessions, .. } = &mut self.mode {
128 *max_sessions = max.max(1);
129 }
130 self
131 }
132
133 pub fn memory(&self) -> SharedMemory {
138 match &self.mode {
139 HistoryMode::Shared(m) => m.clone(),
140 HistoryMode::Sessions { default, .. } => default.clone(),
141 }
142 }
143
144 async fn select_memory(
146 &self,
147 config: &Option<RunnableConfig>,
148 ) -> Result<SharedMemory, LcelError> {
149 match &self.mode {
150 HistoryMode::Shared(m) => Ok(m.clone()),
151 HistoryMode::Sessions {
152 factory,
153 cache,
154 max_sessions,
155 ..
156 } => {
157 let session_id = config
158 .as_ref()
159 .and_then(|c| c.configurable_value("session_id"))
160 .and_then(|v| v.as_str())
161 .ok_or_else(|| {
162 LcelError::Chain(
163 "RunnableWithMessageHistory (session mode) is missing configurable.session_id"
164 .to_string(),
165 )
166 })?;
167 let mut cache = cache.lock().await;
168 if let Some(memory) = cache.slots.get(session_id) {
169 return Ok(memory.clone());
170 }
171 if cache.slots.len() >= *max_sessions {
173 if let Some(oldest) = cache.order.pop_front() {
174 cache.slots.remove(&oldest);
175 }
176 }
177 let memory = factory(session_id);
178 cache.slots.insert(session_id.to_string(), memory.clone());
179 cache.order.push_back(session_id.to_string());
180 Ok(memory)
181 }
182 }
183 }
184}
185
186#[async_trait]
187impl<L> Runnable<String, LLMResult> for RunnableWithMessageHistory<L>
188where
189 L: BaseChatModel + 'static,
190 L::Error: Into<LcelError>,
191{
192 type Error = LcelError;
193
194 async fn invoke(
195 &self,
196 input: String,
197 config: Option<RunnableConfig>,
198 ) -> Result<LLMResult, LcelError> {
199 let memory = self.select_memory(&config).await?;
201
202 let mut memory = memory.lock().await;
206
207 let mut messages = {
209 let vars = memory
210 .load_memory_variables(&HashMap::new())
211 .await
212 .map_err(|e| LcelError::Chain(format!("load memory: {e}")))?;
213 memory_variables_to_messages(&vars)
214 };
215
216 messages.push(Message::human(&input));
218
219 let result = self.llm.chat(messages, config).await.map_err(Into::into)?;
221
222 let inputs = HashMap::from([("input".to_string(), input)]);
225 let outputs = HashMap::from([("output".to_string(), result.content.clone())]);
226 if let Err(e) = memory.save_context(&inputs, &outputs).await {
227 log::warn!("memory save failed (model answer still returned): {e}");
228 }
229
230 Ok(result)
231 }
232}
233
234#[cfg(test)]
235mod tests {
236 use super::*;
237 use futures_util::Stream;
238 use lc_core::language_models::{BaseChatModel, BaseLanguageModel};
239 use lc_core::runnables::RunnableConfig;
240 use lc_schema::MessageType;
241 use serde_json::json;
242 use std::pin::Pin;
243 use std::sync::Mutex as StdMutex;
244
245 struct TestChatModel {
247 seen: Arc<StdMutex<Vec<Vec<Message>>>>,
248 }
249
250 #[async_trait]
251 impl Runnable<Vec<Message>, LLMResult> for TestChatModel {
252 type Error = LcelError;
253
254 async fn invoke(
255 &self,
256 input: Vec<Message>,
257 _config: Option<RunnableConfig>,
258 ) -> Result<LLMResult, LcelError> {
259 self.seen.lock().unwrap().push(input.clone());
260 let last = input.last().map(|m| m.content.clone()).unwrap_or_default();
261 Ok(LLMResult {
262 content: format!("reply to: {last}"),
263 ..Default::default()
264 })
265 }
266 }
267
268 #[async_trait]
269 impl BaseLanguageModel<Vec<Message>, LLMResult> for TestChatModel {
270 fn model_name(&self) -> &str {
271 "test-llm"
272 }
273
274 fn get_num_tokens(&self, text: &str) -> usize {
275 text.len()
276 }
277
278 fn with_temperature(self, _temp: f32) -> Self
279 where
280 Self: Sized,
281 {
282 self
283 }
284
285 fn with_max_tokens(self, _max: usize) -> Self
286 where
287 Self: Sized,
288 {
289 self
290 }
291 }
292
293 #[async_trait]
294 impl BaseChatModel for TestChatModel {
295 async fn chat(
296 &self,
297 messages: Vec<Message>,
298 config: Option<RunnableConfig>,
299 ) -> Result<LLMResult, LcelError> {
300 self.invoke(messages, config).await
301 }
302
303 async fn stream_chat(
304 &self,
305 _messages: Vec<Message>,
306 _config: Option<RunnableConfig>,
307 ) -> Result<Pin<Box<dyn Stream<Item = Result<String, LcelError>> + Send>>, LcelError>
308 {
309 unimplemented!("stream_chat not needed for tests")
310 }
311 }
312
313 fn session_factory(_session_id: &str) -> SharedMemory {
315 Arc::new(Mutex::new(
316 Box::new(ConversationBufferMemory::new().with_return_messages(true))
317 as Box<dyn BaseMemory>,
318 ))
319 }
320
321 #[tokio::test]
323 async fn reads_memory_writes_back_round_trip() {
324 let seen = Arc::new(StdMutex::new(Vec::new()));
326 let llm = TestChatModel { seen: seen.clone() };
327 let memory = ConversationBufferMemory::new().with_return_messages(true);
329 let pipe = RunnableWithMessageHistory::new(llm, memory);
330
331 let r1 = pipe.invoke("我叫什么名字".to_string(), None).await.unwrap();
333 assert_eq!(r1.content, "reply to: 我叫什么名字");
334
335 let r2 = pipe.invoke("再问一次".to_string(), None).await.unwrap();
337 assert_eq!(r2.content, "reply to: 再问一次");
338
339 let calls = seen.lock().unwrap();
341 assert_eq!(calls.len(), 2, "model should be called twice");
342 assert_eq!(calls[0].len(), 1);
344 assert_eq!(calls[0][0].content, "我叫什么名字");
345 assert_eq!(calls[1].len(), 3);
347 assert_eq!(calls[1][0].content, "我叫什么名字");
348 assert!(matches!(calls[1][0].message_type, MessageType::Human));
349 assert!(matches!(calls[1][1].message_type, MessageType::AI));
350 assert_eq!(calls[1][2].content, "再问一次");
351 }
352
353 #[tokio::test]
355 async fn memory_accumulates_across_invocations() {
356 let seen = Arc::new(StdMutex::new(Vec::new()));
358 let llm = TestChatModel { seen: seen.clone() };
359 let memory = ConversationBufferMemory::new().with_return_messages(true);
360 let pipe = RunnableWithMessageHistory::new(llm, memory);
361
362 for turn in ["你好", "你在吗", "再见"] {
364 pipe.invoke(turn.to_string(), None).await.unwrap();
365 }
366
367 let calls = seen.lock().unwrap();
369 assert_eq!(calls.len(), 3);
370 assert_eq!(calls[2].len(), 5); assert_eq!(calls[2][4].content, "再见");
372 }
373
374 #[tokio::test]
376 async fn session_history_same_session_shares_memory() {
377 let seen = Arc::new(StdMutex::new(Vec::new()));
379 let llm = TestChatModel { seen: seen.clone() };
380 let pipe = RunnableWithMessageHistory::with_session_history(llm, session_factory);
381 let cfg = RunnableConfig::new().with_configurable("session_id", json!("s1"));
382
383 let r1 = pipe
385 .invoke("我叫什么名字".to_string(), Some(cfg.clone()))
386 .await
387 .unwrap();
388 assert_eq!(r1.content, "reply to: 我叫什么名字");
389 let r2 = pipe
390 .invoke("再问一次".to_string(), Some(cfg))
391 .await
392 .unwrap();
393 assert_eq!(r2.content, "reply to: 再问一次");
394
395 let calls = seen.lock().unwrap();
397 assert_eq!(calls.len(), 2);
398 assert_eq!(calls[1].len(), 3);
399 assert_eq!(calls[1][0].content, "我叫什么名字");
400 assert!(matches!(calls[1][1].message_type, MessageType::AI));
401 }
402
403 #[tokio::test]
405 async fn session_history_different_sessions_isolated() {
406 let seen = Arc::new(StdMutex::new(Vec::new()));
408 let llm = TestChatModel { seen: seen.clone() };
409 let pipe = RunnableWithMessageHistory::with_session_history(llm, session_factory);
410 let cfg_s1 = RunnableConfig::new().with_configurable("session_id", json!("s1"));
411 let cfg_s2 = RunnableConfig::new().with_configurable("session_id", json!("s2"));
412
413 pipe.invoke("我是 s1".to_string(), Some(cfg_s1.clone()))
415 .await
416 .unwrap();
417 pipe.invoke("还在 s1".to_string(), Some(cfg_s1))
418 .await
419 .unwrap();
420 pipe.invoke("我是 s2".to_string(), Some(cfg_s2))
421 .await
422 .unwrap();
423
424 let calls = seen.lock().unwrap();
426 assert_eq!(calls.len(), 3);
427 assert_eq!(calls[1].len(), 3, "s1 second turn should see history");
428 assert_eq!(calls[2].len(), 1, "s2 first turn should have no history");
429 }
430
431 #[tokio::test]
433 async fn session_history_missing_session_id_errors() {
434 let seen = Arc::new(StdMutex::new(Vec::new()));
436 let llm = TestChatModel { seen };
437 let pipe = RunnableWithMessageHistory::with_session_history(llm, session_factory);
438
439 let err = pipe.invoke("你好".to_string(), None).await.unwrap_err();
441 assert!(matches!(err, LcelError::Chain(_)));
442
443 let cfg_no_sid = RunnableConfig::new();
444 let err = pipe
445 .invoke("你好".to_string(), Some(cfg_no_sid))
446 .await
447 .unwrap_err();
448 assert!(matches!(err, LcelError::Chain(_)));
449 }
450
451 struct BlockingChatModel {
454 seen: Arc<StdMutex<Vec<Vec<Message>>>>,
455 entered: Arc<tokio::sync::Notify>,
456 release: Arc<tokio::sync::Notify>,
457 blocked: Arc<StdMutex<bool>>,
458 }
459
460 #[async_trait]
461 impl Runnable<Vec<Message>, LLMResult> for BlockingChatModel {
462 type Error = LcelError;
463
464 async fn invoke(
465 &self,
466 input: Vec<Message>,
467 _config: Option<RunnableConfig>,
468 ) -> Result<LLMResult, LcelError> {
469 self.seen.lock().unwrap().push(input.clone());
470 let should_block = {
471 let mut blocked = self.blocked.lock().unwrap();
472 if !*blocked {
473 *blocked = true;
474 true
475 } else {
476 false
477 }
478 };
479 self.entered.notify_one();
480 if should_block {
481 self.release.notified().await;
482 }
483 let last = input.last().map(|m| m.content.clone()).unwrap_or_default();
484 Ok(LLMResult {
485 content: format!("reply to: {last}"),
486 ..Default::default()
487 })
488 }
489 }
490
491 #[async_trait]
492 impl BaseLanguageModel<Vec<Message>, LLMResult> for BlockingChatModel {
493 fn model_name(&self) -> &str {
494 "blocking-test-llm"
495 }
496
497 fn get_num_tokens(&self, text: &str) -> usize {
498 text.len()
499 }
500
501 fn with_temperature(self, _temp: f32) -> Self
502 where
503 Self: Sized,
504 {
505 self
506 }
507
508 fn with_max_tokens(self, _max: usize) -> Self
509 where
510 Self: Sized,
511 {
512 self
513 }
514 }
515
516 #[async_trait]
517 impl BaseChatModel for BlockingChatModel {
518 async fn chat(
519 &self,
520 messages: Vec<Message>,
521 config: Option<RunnableConfig>,
522 ) -> Result<LLMResult, LcelError> {
523 self.invoke(messages, config).await
524 }
525
526 async fn stream_chat(
527 &self,
528 _messages: Vec<Message>,
529 _config: Option<RunnableConfig>,
530 ) -> Result<Pin<Box<dyn Stream<Item = Result<String, LcelError>> + Send>>, LcelError>
531 {
532 unimplemented!("stream_chat not needed for tests")
533 }
534 }
535
536 #[tokio::test]
538 async fn session_cache_evicts_oldest_when_over_capacity() {
539 let seen = Arc::new(StdMutex::new(Vec::new()));
541 let llm = TestChatModel { seen: seen.clone() };
542 let pipe = RunnableWithMessageHistory::with_session_history(llm, session_factory)
543 .with_max_sessions(2);
544
545 let cfg_s1 = RunnableConfig::new().with_configurable("session_id", json!("s1"));
546 let cfg_s2 = RunnableConfig::new().with_configurable("session_id", json!("s2"));
547 let cfg_s3 = RunnableConfig::new().with_configurable("session_id", json!("s3"));
548
549 pipe.invoke("s1-turn1".to_string(), Some(cfg_s1.clone()))
551 .await
552 .unwrap();
553 pipe.invoke("s2-turn1".to_string(), Some(cfg_s2))
554 .await
555 .unwrap();
556 pipe.invoke("s3-turn1".to_string(), Some(cfg_s3))
557 .await
558 .unwrap();
559 pipe.invoke("s1-turn2".to_string(), Some(cfg_s1))
561 .await
562 .unwrap();
563
564 let calls = seen.lock().unwrap();
566 assert_eq!(calls.len(), 4);
567 assert_eq!(
568 calls[3].len(),
569 1,
570 "M2a: re-entering s1 after eviction should be a fresh session (no history)"
571 );
572 }
573
574 #[tokio::test]
577 async fn concurrent_same_session_invokes_do_not_lose_history() {
578 let seen = Arc::new(StdMutex::new(Vec::new()));
580 let entered = Arc::new(tokio::sync::Notify::new());
581 let release = Arc::new(tokio::sync::Notify::new());
582 let blocked = Arc::new(StdMutex::new(false));
583 let llm = BlockingChatModel {
584 seen: seen.clone(),
585 entered: entered.clone(),
586 release: release.clone(),
587 blocked,
588 };
589 let pipe = Arc::new(RunnableWithMessageHistory::new(
590 llm,
591 ConversationBufferMemory::new().with_return_messages(true),
592 ));
593
594 let p1 = pipe.clone();
596 let h1 = tokio::spawn(async move { p1.invoke("第一轮".to_string(), None).await });
597 entered.notified().await;
598
599 let p2 = pipe.clone();
601 let h2 = tokio::spawn(async move { p2.invoke("第二轮".to_string(), None).await });
602
603 release.notify_one();
605 let r1 = h1.await.unwrap().unwrap();
606 let r2 = h2.await.unwrap().unwrap();
607 assert_eq!(r1.content, "reply to: 第一轮");
608 assert_eq!(r2.content, "reply to: 第二轮");
609
610 let calls = seen.lock().unwrap();
612 assert_eq!(calls.len(), 2);
613 assert_eq!(
614 calls[1].len(),
615 3,
616 "M2b: second turn should see the full first-turn conversation (user+ai+user), not empty history"
617 );
618 assert_eq!(calls[1][0].content, "第一轮");
619 assert_eq!(calls[1][2].content, "第二轮");
620 }
621}