codekingpro/portable-devtools
114k
1from typing import Any, Dict, List, Type, Union2 3from langchain_core.language_models import BaseLanguageModel4from langchain_core.messages import BaseMessage, SystemMessage, get_buffer_string5from langchain_core.prompts import BasePromptTemplate6from pydantic import Field7 8from langchain_community.graphs import NetworkxEntityGraph9from langchain_community.graphs.networkx_graph import (10 KnowledgeTriple,11 get_entities,12 parse_triples,13)14 15try:16 from langchain_classic.chains.llm import LLMChain17 from langchain_classic.memory.chat_memory import BaseChatMemory18 from langchain_classic.memory.prompt import (19 ENTITY_EXTRACTION_PROMPT,20 KNOWLEDGE_TRIPLE_EXTRACTION_PROMPT,21 )22 from langchain_classic.memory.utils import get_prompt_input_key23 24 class ConversationKGMemory(BaseChatMemory):25 """Knowledge graph conversation memory.26 27 Integrates with external knowledge graph to store and retrieve28 information about knowledge triples in the conversation.29 """30 31 k: int = 232 human_prefix: str = "Human"33 ai_prefix: str = "AI"34 kg: NetworkxEntityGraph = Field(default_factory=NetworkxEntityGraph)35 knowledge_extraction_prompt: BasePromptTemplate = (36 KNOWLEDGE_TRIPLE_EXTRACTION_PROMPT37 )38 entity_extraction_prompt: BasePromptTemplate = ENTITY_EXTRACTION_PROMPT39 llm: BaseLanguageModel40 summary_message_cls: Type[BaseMessage] = SystemMessage41 """Number of previous utterances to include in the context."""42 memory_key: str = "history" #: :meta private:43 44 def load_memory_variables(self, inputs: Dict[str, Any]) -> Dict[str, Any]:45 """Return history buffer."""46 entities = self._get_current_entities(inputs)47 48 summary_strings = []49 for entity in entities:50 knowledge = self.kg.get_entity_knowledge(entity)51 if knowledge:52 summary = f"On {entity}: {'. '.join(knowledge)}."53 summary_strings.append(summary)54 context: Union[str, List]55 if not summary_strings:56 context = [] if self.return_messages else ""57 elif self.return_messages:58 context = [59 self.summary_message_cls(content=text) for text in summary_strings60 ]61 else:62 context = "\n".join(summary_strings)63 64 return {self.memory_key: context}65 66 @property67 def memory_variables(self) -> List[str]:68 """Will always return list of memory variables.69 70 :meta private:71 """72 return [self.memory_key]73 74 def _get_prompt_input_key(self, inputs: Dict[str, Any]) -> str:75 """Get the input key for the prompt."""76 if self.input_key is None:77 return get_prompt_input_key(inputs, self.memory_variables)78 return self.input_key79 80 def _get_prompt_output_key(self, outputs: Dict[str, Any]) -> str:81 """Get the output key for the prompt."""82 if self.output_key is None:83 if len(outputs) != 1:84 raise ValueError(f"One output key expected, got {outputs.keys()}")85 return list(outputs.keys())[0]86 return self.output_key87 88 def get_current_entities(self, input_string: str) -> List[str]:89 chain = LLMChain(llm=self.llm, prompt=self.entity_extraction_prompt)90 buffer_string = get_buffer_string(91 self.chat_memory.messages[-self.k * 2 :],92 human_prefix=self.human_prefix,93 ai_prefix=self.ai_prefix,94 )95 output = chain.predict(96 history=buffer_string,97 input=input_string,98 )99 return get_entities(output)100 101 def _get_current_entities(self, inputs: Dict[str, Any]) -> List[str]:102 """Get the current entities in the conversation."""103 prompt_input_key = self._get_prompt_input_key(inputs)104 return self.get_current_entities(inputs[prompt_input_key])105 106 def get_knowledge_triplets(self, input_string: str) -> List[KnowledgeTriple]:107 chain = LLMChain(llm=self.llm, prompt=self.knowledge_extraction_prompt)108 buffer_string = get_buffer_string(109 self.chat_memory.messages[-self.k * 2 :],110 human_prefix=self.human_prefix,111 ai_prefix=self.ai_prefix,112 )113 output = chain.predict(114 history=buffer_string,115 input=input_string,116 verbose=True,117 )118 knowledge = parse_triples(output)119 return knowledge120 121 def _get_and_update_kg(self, inputs: Dict[str, Any]) -> None:122 """Get and update knowledge graph from the conversation history."""123 prompt_input_key = self._get_prompt_input_key(inputs)124 knowledge = self.get_knowledge_triplets(inputs[prompt_input_key])125 for triple in knowledge:126 self.kg.add_triple(triple)127 128 def save_context(self, inputs: Dict[str, Any], outputs: Dict[str, str]) -> None:129 """Save context from this conversation to buffer."""130 super().save_context(inputs, outputs)131 self._get_and_update_kg(inputs)132 133 def clear(self) -> None:134 """Clear memory contents."""135 super().clear()136 self.kg.clear()137 138except ImportError:139 # Placeholder object140 class ConversationKGMemory: # type: ignore[no-redef]141 pass142 