Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
117 changes: 100 additions & 17 deletions packages/powermem-langchain/src/powermem_langchain/middleware.py
Original file line number Diff line number Diff line change
@@ -1,32 +1,23 @@
"""LangChain middleware entry point for PowerMem.

The VLDB 2026 summer school branch intentionally provides only the public entry
point. Students are expected to replace this placeholder with a LangChain
middleware implementation that satisfies the package contract tests.
"""
"""LangChain middleware entry point for PowerMem."""

from __future__ import annotations

from typing import Any, NotRequired, TypedDict

from langchain.agents.middleware import AgentMiddleware, AgentState
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage


class PowerMemState(AgentState):
"""State schema reserved for the PowerMem middleware implementation."""

powermem_context: NotRequired[str]


class PowerMemStateUpdate(TypedDict):
"""State update returned by memory-loading middleware hooks."""

powermem_context: str
powermem_context: NotRequired[str]
messages: NotRequired[list[BaseMessage]]


class PowerMemMiddleware(AgentMiddleware[PowerMemState, Any, Any]):
"""Placeholder for the summer school implementation."""

state_schema = PowerMemState

def __init__(
Expand All @@ -38,17 +29,109 @@ def __init__(
save_interactions: bool = True,
**kwargs: Any,
) -> None:
pass
super().__init__(**kwargs)
self.memory = memory
self.user_id = user_id or "default"
self.search_limit = search_limit
self.save_interactions = save_interactions

def before_agent(self, state: PowerMemState, runtime) -> PowerMemStateUpdate | None:
pass
latest_user_message = self._latest_message_text(state, HumanMessage)
if not latest_user_message:
return None

try:
result = self.memory.search(
latest_user_message,
user_id=self.user_id,
limit=self.search_limit,
)
except Exception:
return None

memories = self._extract_memory_texts(result)
if not memories:
return None

context = "Relevant long-term memories:\n" + "\n".join(
f"- {memory}" for memory in memories
)

return {
"powermem_context": context,
"messages": [SystemMessage(content=context)],
}

async def abefore_agent(
self,
state: PowerMemState,
runtime,
) -> PowerMemStateUpdate | None:
pass
return self.before_agent(state, runtime)

def after_agent(self, state: PowerMemState, runtime) -> None:
pass
if not self.save_interactions:
return

user_text = self._latest_message_text(state, HumanMessage)
assistant_text = self._latest_message_text(state, AIMessage)

if not user_text or not assistant_text:
return

interaction = f"User: {user_text}\nAssistant: {assistant_text}"

try:
self.memory.add(interaction, user_id=self.user_id, infer=False)
except Exception:
return

@classmethod
def _latest_message_text(
cls,
state: PowerMemState,
message_type: type[BaseMessage],
) -> str:
for message in reversed(state.get("messages", [])):
if isinstance(message, message_type):
return cls._content_to_text(message.content)
return ""

@staticmethod
def _content_to_text(content: Any) -> str:
if isinstance(content, str):
return content

if isinstance(content, list):
parts = []
for item in content:
if isinstance(item, str):
parts.append(item)
elif isinstance(item, dict):
parts.append(str(item.get("text") or item.get("content") or item))
else:
parts.append(str(item))
return "\n".join(parts)

return str(content)

@classmethod
def _extract_memory_texts(cls, result: Any) -> list[str]:
if isinstance(result, dict):
items = result.get("results", [])
elif isinstance(result, list):
items = result
else:
items = []

memories: list[str] = []
for item in items:
if isinstance(item, dict):
text = item.get("memory") or item.get("content") or item.get("text")
else:
text = getattr(item, "memory", None) or getattr(item, "content", None)

if text:
memories.append(cls._content_to_text(text))

return memories
Loading