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(
120 ConversationBufferMemory::new(),
121 ))),
122 },
123 }
124 }
125
126 pub fn with_max_sessions(mut self, max: usize) -> Self {
129 if let HistoryMode::Sessions { max_sessions, .. } = &mut self.mode {
130 *max_sessions = max.max(1);
131 }
132 self
133 }
134
135 pub fn memory(&self) -> SharedMemory {
140 match &self.mode {
141 HistoryMode::Shared(m) => m.clone(),
142 HistoryMode::Sessions { default, .. } => default.clone(),
143 }
144 }
145
146 async fn select_memory(
148 &self,
149 config: &Option<RunnableConfig>,
150 ) -> Result<SharedMemory, LcelError> {
151 match &self.mode {
152 HistoryMode::Shared(m) => Ok(m.clone()),
153 HistoryMode::Sessions {
154 factory,
155 cache,
156 max_sessions,
157 ..
158 } => {
159 let session_id = config
160 .as_ref()
161 .and_then(|c| c.configurable_value("session_id"))
162 .and_then(|v| v.as_str())
163 .ok_or_else(|| {
164 LcelError::Chain(
165 "RunnableWithMessageHistory(session mode) 缺少 configurable.session_id"
166 .to_string(),
167 )
168 })?;
169 let mut cache = cache.lock().await;
170 if let Some(memory) = cache.slots.get(session_id) {
171 return Ok(memory.clone());
172 }
173 if cache.slots.len() >= *max_sessions {
175 if let Some(oldest) = cache.order.pop_front() {
176 cache.slots.remove(&oldest);
177 }
178 }
179 let memory = factory(session_id);
180 cache
181 .slots
182 .insert(session_id.to_string(), memory.clone());
183 cache.order.push_back(session_id.to_string());
184 Ok(memory)
185 }
186 }
187 }
188}
189
190#[async_trait]
191impl<L> Runnable<String, LLMResult> for RunnableWithMessageHistory<L>
192where
193 L: BaseChatModel + 'static,
194 L::Error: Into<LcelError>,
195{
196 type Error = LcelError;
197
198 async fn invoke(
199 &self,
200 input: String,
201 config: Option<RunnableConfig>,
202 ) -> Result<LLMResult, LcelError> {
203 let memory = self.select_memory(&config).await?;
205
206 let mut memory = memory.lock().await;
210
211 let mut messages = {
213 let vars = memory
214 .load_memory_variables(&HashMap::new())
215 .await
216 .map_err(|e| LcelError::Chain(format!("load memory: {e}")))?;
217 memory_variables_to_messages(&vars)
218 };
219
220 messages.push(Message::human(&input));
222
223 let result = self
225 .llm
226 .chat(messages, config)
227 .await
228 .map_err(Into::into)?;
229
230 let inputs = HashMap::from([("input".to_string(), input)]);
233 let outputs = HashMap::from([("output".to_string(), result.content.clone())]);
234 if let Err(e) = memory.save_context(&inputs, &outputs).await {
235 log::warn!("记忆写回失败(模型答案仍照常返回): {e}");
236 }
237
238 Ok(result)
239 }
240}
241
242#[cfg(test)]
243mod tests {
244 use super::*;
245 use futures_util::Stream;
246 use lc_core::language_models::{BaseChatModel, BaseLanguageModel};
247 use lc_core::runnables::RunnableConfig;
248 use lc_schema::MessageType;
249 use serde_json::json;
250 use std::pin::Pin;
251 use std::sync::Mutex as StdMutex;
252
253 struct TestChatModel {
255 seen: Arc<StdMutex<Vec<Vec<Message>>>>,
256 }
257
258 #[async_trait]
259 impl Runnable<Vec<Message>, LLMResult> for TestChatModel {
260 type Error = LcelError;
261
262 async fn invoke(
263 &self,
264 input: Vec<Message>,
265 _config: Option<RunnableConfig>,
266 ) -> Result<LLMResult, LcelError> {
267 self.seen.lock().unwrap().push(input.clone());
268 let last = input
269 .last()
270 .map(|m| m.content.clone())
271 .unwrap_or_default();
272 Ok(LLMResult {
273 content: format!("reply to: {last}"),
274 ..Default::default()
275 })
276 }
277 }
278
279 #[async_trait]
280 impl BaseLanguageModel<Vec<Message>, LLMResult> for TestChatModel {
281 fn model_name(&self) -> &str {
282 "test-llm"
283 }
284
285 fn get_num_tokens(&self, text: &str) -> usize {
286 text.len()
287 }
288
289 fn with_temperature(self, _temp: f32) -> Self
290 where
291 Self: Sized,
292 {
293 self
294 }
295
296 fn with_max_tokens(self, _max: usize) -> Self
297 where
298 Self: Sized,
299 {
300 self
301 }
302 }
303
304 #[async_trait]
305 impl BaseChatModel for TestChatModel {
306 async fn chat(
307 &self,
308 messages: Vec<Message>,
309 config: Option<RunnableConfig>,
310 ) -> Result<LLMResult, LcelError> {
311 self.invoke(messages, config).await
312 }
313
314 async fn stream_chat(
315 &self,
316 _messages: Vec<Message>,
317 _config: Option<RunnableConfig>,
318 ) -> Result<
319 Pin<Box<dyn Stream<Item = Result<String, LcelError>> + Send>>,
320 LcelError,
321 > {
322 unimplemented!("stream_chat not needed for tests")
323 }
324 }
325
326 fn session_factory(
328 _session_id: &str,
329 ) -> SharedMemory {
330 Arc::new(Mutex::new(Box::new(
331 ConversationBufferMemory::new().with_return_messages(true),
332 ) as Box<dyn BaseMemory>))
333 }
334
335 #[tokio::test]
337 async fn reads_memory_writes_back_round_trip() {
338 let seen = Arc::new(StdMutex::new(Vec::new()));
340 let llm = TestChatModel { seen: seen.clone() };
341 let memory = ConversationBufferMemory::new().with_return_messages(true);
343 let pipe = RunnableWithMessageHistory::new(llm, memory);
344
345 let r1 = pipe.invoke("我叫什么名字".to_string(), None).await.unwrap();
347 assert_eq!(r1.content, "reply to: 我叫什么名字");
348
349 let r2 = pipe.invoke("再问一次".to_string(), None).await.unwrap();
351 assert_eq!(r2.content, "reply to: 再问一次");
352
353 let calls = seen.lock().unwrap();
355 assert_eq!(calls.len(), 2, "应调用模型两次");
356 assert_eq!(calls[0].len(), 1);
358 assert_eq!(calls[0][0].content, "我叫什么名字");
359 assert_eq!(calls[1].len(), 3);
361 assert_eq!(calls[1][0].content, "我叫什么名字");
362 assert!(matches!(calls[1][0].message_type, MessageType::Human));
363 assert!(matches!(calls[1][1].message_type, MessageType::AI));
364 assert_eq!(calls[1][2].content, "再问一次");
365 }
366
367 #[tokio::test]
369 async fn memory_accumulates_across_invocations() {
370 let seen = Arc::new(StdMutex::new(Vec::new()));
372 let llm = TestChatModel { seen: seen.clone() };
373 let memory = ConversationBufferMemory::new().with_return_messages(true);
374 let pipe = RunnableWithMessageHistory::new(llm, memory);
375
376 for turn in ["你好", "你在吗", "再见"] {
378 pipe.invoke(turn.to_string(), None).await.unwrap();
379 }
380
381 let calls = seen.lock().unwrap();
383 assert_eq!(calls.len(), 3);
384 assert_eq!(calls[2].len(), 5); assert_eq!(calls[2][4].content, "再见");
386 }
387
388 #[tokio::test]
390 async fn session_history_same_session_shares_memory() {
391 let seen = Arc::new(StdMutex::new(Vec::new()));
393 let llm = TestChatModel { seen: seen.clone() };
394 let pipe = RunnableWithMessageHistory::with_session_history(llm, session_factory);
395 let cfg = RunnableConfig::new().with_configurable("session_id", json!("s1"));
396
397 let r1 = pipe.invoke("我叫什么名字".to_string(), Some(cfg.clone())).await.unwrap();
399 assert_eq!(r1.content, "reply to: 我叫什么名字");
400 let r2 = pipe.invoke("再问一次".to_string(), Some(cfg)).await.unwrap();
401 assert_eq!(r2.content, "reply to: 再问一次");
402
403 let calls = seen.lock().unwrap();
405 assert_eq!(calls.len(), 2);
406 assert_eq!(calls[1].len(), 3);
407 assert_eq!(calls[1][0].content, "我叫什么名字");
408 assert!(matches!(calls[1][1].message_type, MessageType::AI));
409 }
410
411 #[tokio::test]
413 async fn session_history_different_sessions_isolated() {
414 let seen = Arc::new(StdMutex::new(Vec::new()));
416 let llm = TestChatModel { seen: seen.clone() };
417 let pipe = RunnableWithMessageHistory::with_session_history(llm, session_factory);
418 let cfg_s1 = RunnableConfig::new().with_configurable("session_id", json!("s1"));
419 let cfg_s2 = RunnableConfig::new().with_configurable("session_id", json!("s2"));
420
421 pipe.invoke("我是 s1".to_string(), Some(cfg_s1.clone())).await.unwrap();
423 pipe.invoke("还在 s1".to_string(), Some(cfg_s1)).await.unwrap();
424 pipe.invoke("我是 s2".to_string(), Some(cfg_s2)).await.unwrap();
425
426 let calls = seen.lock().unwrap();
428 assert_eq!(calls.len(), 3);
429 assert_eq!(calls[1].len(), 3, "s1 第二轮应看到历史");
430 assert_eq!(calls[2].len(), 1, "s2 第一轮应无历史");
431 }
432
433 #[tokio::test]
435 async fn session_history_missing_session_id_errors() {
436 let seen = Arc::new(StdMutex::new(Vec::new()));
438 let llm = TestChatModel { seen };
439 let pipe = RunnableWithMessageHistory::with_session_history(llm, session_factory);
440
441 let err = pipe.invoke("你好".to_string(), None).await.unwrap_err();
443 assert!(matches!(err, LcelError::Chain(_)));
444
445 let cfg_no_sid = RunnableConfig::new();
446 let err = pipe
447 .invoke("你好".to_string(), Some(cfg_no_sid))
448 .await
449 .unwrap_err();
450 assert!(matches!(err, LcelError::Chain(_)));
451 }
452
453 struct BlockingChatModel {
456 seen: Arc<StdMutex<Vec<Vec<Message>>>>,
457 entered: Arc<tokio::sync::Notify>,
458 release: Arc<tokio::sync::Notify>,
459 blocked: Arc<StdMutex<bool>>,
460 }
461
462 #[async_trait]
463 impl Runnable<Vec<Message>, LLMResult> for BlockingChatModel {
464 type Error = LcelError;
465
466 async fn invoke(
467 &self,
468 input: Vec<Message>,
469 _config: Option<RunnableConfig>,
470 ) -> Result<LLMResult, LcelError> {
471 self.seen.lock().unwrap().push(input.clone());
472 let should_block = {
473 let mut blocked = self.blocked.lock().unwrap();
474 if !*blocked {
475 *blocked = true;
476 true
477 } else {
478 false
479 }
480 };
481 self.entered.notify_one();
482 if should_block {
483 self.release.notified().await;
484 }
485 let last = input
486 .last()
487 .map(|m| m.content.clone())
488 .unwrap_or_default();
489 Ok(LLMResult {
490 content: format!("reply to: {last}"),
491 ..Default::default()
492 })
493 }
494 }
495
496 #[async_trait]
497 impl BaseLanguageModel<Vec<Message>, LLMResult> for BlockingChatModel {
498 fn model_name(&self) -> &str {
499 "blocking-test-llm"
500 }
501
502 fn get_num_tokens(&self, text: &str) -> usize {
503 text.len()
504 }
505
506 fn with_temperature(self, _temp: f32) -> Self
507 where
508 Self: Sized,
509 {
510 self
511 }
512
513 fn with_max_tokens(self, _max: usize) -> Self
514 where
515 Self: Sized,
516 {
517 self
518 }
519 }
520
521 #[async_trait]
522 impl BaseChatModel for BlockingChatModel {
523 async fn chat(
524 &self,
525 messages: Vec<Message>,
526 config: Option<RunnableConfig>,
527 ) -> Result<LLMResult, LcelError> {
528 self.invoke(messages, config).await
529 }
530
531 async fn stream_chat(
532 &self,
533 _messages: Vec<Message>,
534 _config: Option<RunnableConfig>,
535 ) -> Result<
536 Pin<Box<dyn Stream<Item = Result<String, LcelError>> + Send>>,
537 LcelError,
538 > {
539 unimplemented!("stream_chat not needed for tests")
540 }
541 }
542
543 #[tokio::test]
545 async fn session_cache_evicts_oldest_when_over_capacity() {
546 let seen = Arc::new(StdMutex::new(Vec::new()));
548 let llm = TestChatModel { seen: seen.clone() };
549 let pipe = RunnableWithMessageHistory::with_session_history(llm, session_factory)
550 .with_max_sessions(2);
551
552 let cfg_s1 = RunnableConfig::new().with_configurable("session_id", json!("s1"));
553 let cfg_s2 = RunnableConfig::new().with_configurable("session_id", json!("s2"));
554 let cfg_s3 = RunnableConfig::new().with_configurable("session_id", json!("s3"));
555
556 pipe.invoke("s1-turn1".to_string(), Some(cfg_s1.clone()))
558 .await
559 .unwrap();
560 pipe.invoke("s2-turn1".to_string(), Some(cfg_s2))
561 .await
562 .unwrap();
563 pipe.invoke("s3-turn1".to_string(), Some(cfg_s3))
564 .await
565 .unwrap();
566 pipe.invoke("s1-turn2".to_string(), Some(cfg_s1))
568 .await
569 .unwrap();
570
571 let calls = seen.lock().unwrap();
573 assert_eq!(calls.len(), 4);
574 assert_eq!(
575 calls[3].len(),
576 1,
577 "M2a: s1 槽被淘汰后重入应为全新会话(无历史)"
578 );
579 }
580
581 #[tokio::test]
584 async fn concurrent_same_session_invokes_do_not_lose_history() {
585 let seen = Arc::new(StdMutex::new(Vec::new()));
587 let entered = Arc::new(tokio::sync::Notify::new());
588 let release = Arc::new(tokio::sync::Notify::new());
589 let blocked = Arc::new(StdMutex::new(false));
590 let llm = BlockingChatModel {
591 seen: seen.clone(),
592 entered: entered.clone(),
593 release: release.clone(),
594 blocked,
595 };
596 let pipe = Arc::new(RunnableWithMessageHistory::new(
597 llm,
598 ConversationBufferMemory::new().with_return_messages(true),
599 ));
600
601 let p1 = pipe.clone();
603 let h1 = tokio::spawn(async move { p1.invoke("第一轮".to_string(), None).await });
604 entered.notified().await;
605
606 let p2 = pipe.clone();
608 let h2 = tokio::spawn(async move { p2.invoke("第二轮".to_string(), None).await });
609
610 release.notify_one();
612 let r1 = h1.await.unwrap().unwrap();
613 let r2 = h2.await.unwrap().unwrap();
614 assert_eq!(r1.content, "reply to: 第一轮");
615 assert_eq!(r2.content, "reply to: 第二轮");
616
617 let calls = seen.lock().unwrap();
619 assert_eq!(calls.len(), 2);
620 assert_eq!(
621 calls[1].len(),
622 3,
623 "M2b: 第二轮应看到第一轮完整对话(user+ai+user),而非空历史"
624 );
625 assert_eq!(calls[1][0].content, "第一轮");
626 assert_eq!(calls[1][2].content, "第二轮");
627 }
628}