File size: 5,640 Bytes
c0ff71f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 | from dotenv import load_dotenv
from langchain_openai import ChatOpenAI
from pydantic import BaseModel, Field
from trustcall import create_extractor
from langchain_core.messages import HumanMessage, SystemMessage, AIMessage
from typing import TypedDict, Literal
load_dotenv()
model = ChatOpenAI(model="gpt-4.1-mini", temperature=0)
class Memory(BaseModel):
content: str = Field(description="The main content of the memory. For example: User expressed interest in learning about French.")
class MemoryCollection(BaseModel):
memories: list[Memory] = Field(description="A list of memories about the user.")
# Inspect the tool calls made by Trustcall
class Spy:
def __init__(self):
self.called_tools = []
def __call__(self, run):
# Collect information about the tool calls made by the extractor.
q = [run]
while q:
r = q.pop()
if r.child_runs:
q.extend(r.child_runs)
if r.run_type == "chat_model":
self.called_tools.append(
r.outputs["generations"][0][0]["message"]["kwargs"]["tool_calls"]
)
# Initialize the spy
spy = Spy()
# Create the extractor
trustcall_extractor = create_extractor(
model,
tools=[Memory],
tool_choice="Memory",
enable_inserts=True,
)
# Add the spy as a listener
trustcall_extractor_see_all_tool_calls = trustcall_extractor.with_listeners(on_end=spy)
# Instruction
instruction = """Extract memories from the following conversation:"""
# Conversation
conversation = [HumanMessage(content="Hi, I'm Lance."),
AIMessage(content="Nice to meet you, Lance."),
HumanMessage(content="This morning I had a nice bike ride in San Francisco.")]
# Invoke the extractor
result = trustcall_extractor.invoke({"messages": [SystemMessage(content=instruction)] + conversation})
print("------------------")
print("Mensaje: 1")
print("------------------")
# Messages contain the tool calls
for m in result["messages"]:
m.pretty_print()
# Update the conversation
updated_conversation = [AIMessage(content="That's great, did you do after?"),
HumanMessage(content="I went to Tartine and ate a croissant."),
AIMessage(content="What else is on your mind?"),
HumanMessage(content="I was thinking about my Japan, and going back this winter!"),]
print("------------------")
print("Mensaje: 2: Update system message")
print("------------------")
# Update the instruction
system_msg = """Update existing memories and create new ones based on the following conversation:"""
# We'll save existing memories, giving them an ID, key (tool name), and value
tool_name = "Memory"
existing_memories = [(str(i), tool_name, memory.model_dump()) for i, memory in enumerate(result["responses"])] if result["responses"] else None
print(existing_memories)
# Invoke the extractor with our updated conversation and existing memories
result = trustcall_extractor_see_all_tool_calls.invoke({"messages": updated_conversation,
"existing": existing_memories})
print("------------------")
print("Mensaje: 3: metadata and tool calls")
print("------------------")
# Metadata contains the tool call
for m in result["response_metadata"]:
print(m)
print("------------------")
print("Mensaje: 4: metadata and tool calls")
print("------------------")
# Messages contain the tool calls
for m in result["messages"]:
m.pretty_print()
print("------------------")
print("Mensaje: 5: Parsed responses")
print("------------------")
# Parsed responses
for m in result["responses"]:
print(m)
print("------------------")
print("Mensaje: 6: Inspect the tool calls made by Trustcall")
print("------------------")
# Inspect the tool calls made by Trustcall
print(spy.called_tools)
def extract_tool_info(tool_calls, schema_name="Memory"):
"""Extract information from tool calls for both patches and new memories.
Args:
tool_calls: List of tool calls from the model
schema_name: Name of the schema tool (e.g., "Memory", "ToDo", "Profile")
"""
# Initialize list of changes
changes = []
for call_group in tool_calls:
for call in call_group:
if call['name'] == 'PatchDoc':
changes.append({
'type': 'update',
'doc_id': call['args']['json_doc_id'],
'planned_edits': call['args']['planned_edits'],
'value': call['args']['patches'][0]['value']
})
elif call['name'] == schema_name:
changes.append({
'type': 'new',
'value': call['args']
})
# Format results as a single string
result_parts = []
for change in changes:
if change['type'] == 'update':
result_parts.append(
f"Document {change['doc_id']} updated:\n"
f"Plan: {change['planned_edits']}\n"
f"Added content: {change['value']}"
)
else:
result_parts.append(
f"New {schema_name} created:\n"
f"Content: {change['value']}"
)
return "\n\n".join(result_parts)
print("------------------")
print("Mensaje: 7: Extracted changes")
print("------------------")
# Inspect spy.called_tools to see exactly what happened during the extraction
schema_name = "Memory"
changes = extract_tool_info(spy.called_tools, schema_name)
print(changes) |