SpatialAgent_N2_DEEPAGENT / spatialagent /agent /spatialagent_deep.py
jkkkkyuedtrt's picture
Upload 381 files
bd52a47 verified
Raw
History Blame Contribute Delete
30.5 kB
"""
SpatialAgent - Deepagents-based implementation.
Refactored SpatialAgent using the deepagents framework for:
- Built-in task planning (write_todos)
- Filesystem operations
- Sub-agent delegation
- Context management
- Better scalability and maintainability
"""
from typing import Annotated, List, Dict, Any, TypedDict, Literal, Optional
import os, re, uuid, logging
from dataclasses import dataclass, field
from deepagents import create_deep_agent
from deepagents.backends.filesystem import FilesystemBackend
from deepagents.backends.state import StateBackend
from deepagents.middleware.skills import SkillsMiddleware
from ..hooks_deep import HooksMiddleware
from langchain_core.messages import BaseMessage, HumanMessage
from langchain_core.tools import BaseTool
from .make_prompt import AgentPrompts
from .make_llm import make_llm, DEFAULT_CLAUDE_MODEL
from .utils import load_all_tools
@dataclass
class SpatialAgentConfig:
"""Configuration for SpatialAgent using deepagents."""
llm: Any = None
data_path: str = "./data"
save_path: str = "./experiments"
tool_retrieval: bool = False
tool_retrieval_method: str = "llm"
min_tools: int = 5
max_tools: int = 20
skill_retrieval: bool = True
num_skills: int = 1
auto_interpret_figures: bool = True
act_timeout: int = 1800
web_search_model: str = "gemini-3-flash-preview"
backend_type: Literal["filesystem", "state"] = "filesystem"
class SpatialAgent:
"""SpatialAgent refactored using deepagents framework."""
def __init__(self, config: SpatialAgentConfig = None, **kwargs):
"""
Initialize SpatialAgent with deepagents.
Args:
config: SpatialAgentConfig instance with configuration options
**kwargs: Additional configuration options (for backward compatibility)
"""
if config is None:
config = SpatialAgentConfig()
for key, value in kwargs.items():
if hasattr(config, key):
setattr(config, key, value)
self.config = config
if config.llm is None:
print(f"No LLM provided, using default: {DEFAULT_CLAUDE_MODEL}", flush=True)
config.llm = make_llm(DEFAULT_CLAUDE_MODEL)
self.llm = config.llm
self.tool_retrieval = config.tool_retrieval
self.min_tools = config.min_tools
self.max_tools = config.max_tools
self.auto_interpret_figures = config.auto_interpret_figures
self.act_timeout = config.act_timeout
from . import set_agent_model
model_name = None
for attr in ['deployment_name', 'model_id', 'model_name', 'model']:
val = getattr(config.llm, attr, None)
if val and isinstance(val, str):
model_name = val
break
if not model_name:
model_name = "unknown"
set_agent_model(model_name, config.llm)
self.web_search_model = config.web_search_model if config.web_search_model else model_name
self.observation_log = []
self._observation_log_path = os.path.join(config.save_path, "observation_log.jsonl")
self.conversation_history = {}
self.default_thread_id = str(uuid.uuid4())
data_path = os.path.abspath(config.data_path)
save_path = os.path.abspath(config.save_path)
self.save_path = save_path
self.data_path = data_path
print(f"Data path: {data_path}", flush=True)
print(f"Save path: {save_path}", flush=True)
os.makedirs(save_path, exist_ok=True)
os.makedirs(data_path, exist_ok=True)
self._build_backend()
self._build_skills()
self._load_tools()
self._build_system_prompt()
self._build_agent()
def _load_tools(self):
"""Load all spatial transcriptomics tools."""
print("Loading tools from tool modules...", flush=True)
self.tools = load_all_tools(save_path=self.save_path, data_path=self.data_path)
print(f"Loaded {len(self.tools)} tools", flush=True)
if self.config.skill_retrieval and hasattr(self, 'skills_metadata') and self.skills_metadata:
self._add_skill_tool()
def _add_skill_tool(self):
"""Add a use_skill tool to read skill details."""
from langchain_core.tools import tool
skills_meta = self.skills_metadata
@tool
def use_skill(skill_name: str) -> str:
"""Read the full instructions for a spatial transcriptomics analysis skill.
Args:
skill_name: Name of the skill to read (e.g., "liana-analysis", "spatial-deconvolution")
Returns:
Full skill documentation with workflow steps, best practices, and examples.
"""
for skill in skills_meta:
if skill['name'] == skill_name or skill['name'].replace('-', '_') == skill_name:
try:
with open(skill['path'], 'r') as f:
return f.read()
except Exception as e:
return f"Error reading skill: {e}"
available = [s['name'] for s in skills_meta]
return f"Skill '{skill_name}' not found. Available skills: {', '.join(available)}"
self.tools.append(use_skill)
print(f"Added use_skill tool for {len(self.skills_metadata)} skills", flush=True)
def _build_system_prompt(self):
"""Build system prompt with spatial transcriptomics domain knowledge."""
tool_list = []
for tool in self.tools:
tool_name = getattr(tool, 'name', tool.__class__.__name__)
tool_desc = getattr(tool, 'description', 'No description')
tool_list.append((tool_name, tool_desc))
tool_names = [t[0] for t in tool_list]
skills_section = ""
if hasattr(self, 'skills_metadata') and self.skills_metadata:
skill_list_str = "\n".join([f"- **{s['name']}**: {s['description'][:80]}" for s in self.skills_metadata])
skills_section = f"""
## SPECIALIZED SPATIAL TRANSCRIPTOMICS SKILLS
You also have access to {len(self.skills_metadata)} curated workflow skills. Use the `use_skill` tool to read full instructions.
Available skills:
{skill_list_str}
## How to Use Skills
1. When a user's task matches a skill's domain, call `use_skill(skill_name="...")` to get detailed workflow instructions
2. Follow the skill's step-by-step guidance
3. Use the specialized tools listed above to execute each step
## TOOLS vs SKILLS
- **TOOLS (工具)**: Individual callable functions (search_panglao, liana_inference, execute_python, etc.) - 72 total
- **SKILLS (技能)**: Curated workflow guides that tell you WHICH tools to use and WHEN - {len(self.skills_metadata)} total
- Use `use_skill` to read skill details, then use tools to execute
"""
self.system_prompt = f"""You are a helpful assistant specialized in spatial transcriptomics analysis.
## CONVERSATION CONTEXT
IMPORTANT: You HAVE FULL ACCESS to conversation history. ALWAYS read and remember ALL previous messages in the conversation. Use this context to provide personalized responses. You can and should reference information from earlier messages.
## COMPLEX TASK PLANNING
When you receive a complex task that requires multiple steps or tool calls, YOU MUST FIRST OUTPUT A PLAN (todo list) before executing any tools. This helps organize the workflow and ensures all necessary steps are completed.
### HOW TO DETERMINE IF A TASK IS COMPLEX:
- Simple tasks: Can be answered in 1-2 sentences or with a single tool call (e.g., "What is my name?", "Query disease genes for prostate cancer")
- Complex tasks: Require multiple steps, multiple tool calls, or synthesis of information (e.g., "Design a 50-gene panel", "Analyze spatial transcriptomics data", "Generate a comprehensive report")
### PLAN FORMAT (MUST FOLLOW):
For complex tasks, output your plan FIRST in this format:
```plan
## Task: [Task Name]
### Steps:
1. [Step 1 description] - [Tool to use if applicable]
2. [Step 2 description] - [Tool to use if applicable]
3. [Step 3 description] - [Tool to use if applicable]
...
### Expected Output:
[Brief description of what the final output will include]
```
After outputting the plan, you can start executing the steps one by one using tool calls.
## TOOL CALLING FORMAT (MUST FOLLOW)
When you need to call a tool, output ONLY a JSON object with "name" and "arguments" fields. DO NOT use markdown code blocks.
EXACT FORMAT:
{{
"name": "TOOL_NAME",
"arguments": {{
"PARAM1": "VALUE1",
"PARAM2": VALUE2
}}
}}
## CRITICAL RULES
- NEVER use ```json or any markdown code blocks
- NEVER use <act> tags
- ALWAYS use plain JSON format for tool calls
- ALWAYS include both "name" and "arguments" fields
- Parameter names must match exactly (e.g., "disease" not "diseases")
- Execute ONE tool at a time, wait for results, then continue
- ALWAYS remember and use conversation history for context
- NEVER say you don't have access to conversation history - you ALWAYS have access
- FOR SIMPLE QUESTIONS that can be answered from conversation history alone (like "what is my name?", "what did we talk about?"), DO NOT CALL TOOLS. Answer directly using the conversation history.
- FOR COMPLEX TASKS, ALWAYS OUTPUT A PLAN FIRST before executing tools.
## TOOLS AVAILABLE
Database: search_panglao, search_cellmarker2, search_czi_datasets, query_tissue_expression, query_disease_genes
Literature: query_pubmed, search_semantic_scholar, web_search
Analytics: liana_inference, squidpy_ligrec, cellphonedb_analysis, tangram_map_cells, spagcn_clustering, scanpy_score_genes
Interpretation: annotate_cell_types, annotate_tissue_niches
Coding: execute_python, execute_bash
Subagent: report_subagent, verification_subagent
## WHEN TO USE SKILLS
Use `use_skill(skill_name="...")` for complex workflows. For simple queries, call tools directly.
"""
def _build_backend(self):
"""Build the filesystem backend for the agent."""
if self.config.backend_type == "filesystem":
self.backend = FilesystemBackend(
root_dir=self.save_path,
virtual_mode=True,
max_file_size_mb=50
)
print(f"Using FilesystemBackend with root_dir: {self.save_path}", flush=True)
else:
self.backend = StateBackend()
print("Using StateBackend (ephemeral storage)", flush=True)
def _build_skills(self):
"""Load skill templates for common spatial transcriptomics workflows."""
if self.config.skill_retrieval:
skills_source_dir = os.path.join(os.path.dirname(__file__), '..', 'skill')
skills_source_dir = os.path.abspath(skills_source_dir)
skills_target_dir = os.path.join(self.save_path, 'skills')
if os.path.exists(skills_source_dir):
self._copy_skills(skills_source_dir, skills_target_dir)
self.skills_path = skills_target_dir
print(f"Skills loaded from: {skills_source_dir}", flush=True)
print(f"Skills copied to: {skills_target_dir}", flush=True)
self._load_skills_metadata()
else:
self.skills_path = None
self.skills_metadata = []
print("Skills source directory not found, skill retrieval disabled", flush=True)
else:
self.skills_path = None
self.skills_metadata = []
def _load_skills_metadata(self):
"""Load skill metadata from skill directories."""
import re
self.skills_metadata = []
if not self.skills_path or not os.path.exists(self.skills_path):
return
for item in sorted(os.listdir(self.skills_path)):
item_path = os.path.join(self.skills_path, item)
if not os.path.isdir(item_path):
continue
skill_md = os.path.join(item_path, 'SKILL.md')
if not os.path.exists(skill_md):
continue
with open(skill_md, 'r') as f:
content = f.read()
name = item
description = ""
frontmatter_match = re.match(r'^---\n(.*?)\n---', content, re.DOTALL)
if frontmatter_match:
frontmatter = frontmatter_match.group(1)
for line in frontmatter.split('\n'):
if line.startswith('name:'):
name = line.split(':', 1)[1].strip()
elif line.startswith('description:'):
description = line.split(':', 1)[1].strip()
if not description:
lines = content.strip().split('\n')
for line in lines:
if line.startswith('# '):
continue
if line.strip():
description = line.strip()
break
self.skills_metadata.append({
'name': name,
'description': description,
'path': skill_md
})
print(f"Loaded {len(self.skills_metadata)} skills", flush=True)
def _copy_skills(self, source_dir: str, target_dir: str):
"""Copy skill directories from source to target."""
import shutil
os.makedirs(target_dir, exist_ok=True)
for item in os.listdir(source_dir):
item_path = os.path.join(source_dir, item)
if os.path.isdir(item_path):
target_item = os.path.join(target_dir, item)
if os.path.exists(target_item):
shutil.rmtree(target_item)
shutil.copytree(item_path, target_item)
def _build_agent(self):
"""Build the deep agent using langgraph with proper tool binding and memory."""
print("Building deep agent...", flush=True)
from langchain_core.messages import SystemMessage
from langgraph.graph import StateGraph, END
from langgraph.prebuilt import ToolNode
from langgraph.checkpoint.memory import MemorySaver
def should_continue(state):
messages = state["messages"]
last_message = messages[-1]
if hasattr(last_message, 'tool_calls') and last_message.tool_calls:
return "tools"
return END
def call_model(state):
messages = state["messages"]
system_message = SystemMessage(content=self.system_prompt)
full_messages = [system_message] + messages
content_to_send = ""
for msg in full_messages:
if isinstance(msg, SystemMessage):
content_to_send += f"System: {msg.content}\n"
elif hasattr(msg, 'content'):
role = "User" if isinstance(msg, HumanMessage) else "Assistant"
content_to_send += f"{role}: {msg.content}\n"
response = self.llm.invoke(content_to_send)
response_content = response.content if hasattr(response, 'content') else str(response)
import json
import re
import uuid
tool_calls = []
json_match = re.search(r'\{[\s\S]*"name":\s*"[^"]+"[\s\S]*\}', response_content)
if json_match:
try:
tool_call_json = json.loads(json_match.group())
if "name" in tool_call_json and "arguments" in tool_call_json:
tool_calls = [{
"id": str(uuid.uuid4()),
"name": tool_call_json["name"],
"args": tool_call_json["arguments"]
}]
except json.JSONDecodeError:
pass
from langchain_core.messages import AIMessage
if tool_calls:
ai_response = AIMessage(content=response_content, tool_calls=tool_calls)
else:
ai_response = AIMessage(content=response_content)
return {"messages": messages + [ai_response]}
tool_node = ToolNode(self.tools)
workflow = StateGraph(dict)
workflow.add_node("agent", call_model)
workflow.add_node("tools", tool_node)
workflow.set_entry_point("agent")
workflow.add_conditional_edges("agent", should_continue)
workflow.add_edge("tools", "agent")
memory_saver = MemorySaver()
self.agent = workflow.compile(checkpointer=memory_saver)
print("Deep agent built successfully with memory!", flush=True)
def _log_observation(self, step_number: int, code: str, result: str, figure_interpretations: str = ""):
"""Log an observation for later report generation."""
import json
from datetime import datetime
entry = {
"step": step_number,
"timestamp": datetime.now().isoformat(),
"code_snippet": code,
"result_summary": result,
"figure_interpretations": figure_interpretations,
}
self.observation_log.append(entry)
try:
with open(self._observation_log_path, 'a') as f:
f.write(json.dumps(entry) + '\n')
except Exception as e:
print(f"Warning: Could not write to observation log: {e}")
def _display_figures(self, code_context: str = "", user_query: str = "") -> str:
"""Display any new image files and optionally interpret them using vision LLM."""
interpretations = []
try:
from spatialagent.tool.coding import get_new_image_files
image_files = get_new_image_files()
if not image_files:
return ""
try:
from IPython.display import display, Image, SVG
import os
print(f"📊 Displaying {len(image_files)} figure(s)...")
for img_path in image_files:
if not os.path.exists(img_path):
print(f"⚠️ File not found: {img_path}")
continue
ext = os.path.splitext(img_path)[1].lower()
if ext == '.svg':
display(SVG(filename=img_path))
elif ext in ('.png', '.jpg', '.jpeg'):
display(Image(filename=img_path))
elif ext == '.pdf':
print(f"📄 Created: {os.path.basename(img_path)}")
except ImportError:
import os
print(f"[{len(image_files)} figure(s) created: {', '.join(os.path.basename(f) for f in image_files)}]")
if self.auto_interpret_figures and image_files:
print(f"🔍 Interpreting {len(image_files)} figure(s)...")
from spatialagent.tool.interpretation import interpret_figure
for img_path in image_files:
import os
if not os.path.exists(img_path):
continue
ext = os.path.splitext(img_path)[1].lower()
if ext == '.pdf':
continue
try:
context = self._infer_figure_context(code_context, img_path, user_query)
interpretation = interpret_figure.invoke({
"image_path": img_path,
"context": context,
"analysis_focus": "general"
})
fig_name = os.path.basename(img_path)
interpretations.append(f"\n### Figure Interpretation: {fig_name}\n{interpretation}")
except Exception as e:
print(f"⚠️ Could not interpret {os.path.basename(img_path)}: {e}")
except Exception as e:
print(f"⚠️ Error displaying/interpreting figures: {e}")
return "\n".join(interpretations) if interpretations else ""
def _infer_figure_context(self, code: str, img_path: str, user_query: str = "") -> str:
"""Infer the context/type of a figure from the code that generated it."""
import os
fig_name = os.path.basename(img_path)
context_parts = [f"Figure: {fig_name}"]
code_lower = code.lower()
plot_types = []
if "umap" in code_lower:
plot_types.append("UMAP dimensionality reduction")
if "tsne" in code_lower or "t-sne" in code_lower:
plot_types.append("t-SNE dimensionality reduction")
if "pca" in code_lower and "plot" in code_lower:
plot_types.append("PCA plot")
if "sc.pl.spatial" in code_lower or "sq.pl.spatial" in code_lower or "spatial_scatter" in code_lower:
plot_types.append("Spatial plot showing tissue coordinates")
if "heatmap" in code_lower or "sns.heatmap" in code_lower or "clustermap" in code_lower:
plot_types.append("Heatmap visualization")
if "violin" in code_lower:
plot_types.append("Violin plot")
if "dotplot" in code_lower or "dot_plot" in code_lower:
plot_types.append("Dot plot")
if "stacked_violin" in code_lower:
plot_types.append("Stacked violin plot")
if "matrixplot" in code_lower:
plot_types.append("Matrix plot")
if "rank_genes" in code_lower:
plot_types.append("Ranked genes plot")
if "barplot" in code_lower or "bar(" in code_lower or "barh(" in code_lower:
plot_types.append("Bar plot")
if "boxplot" in code_lower:
plot_types.append("Box plot")
if "scatter" in code_lower and "spatial" not in code_lower:
plot_types.append("Scatter plot")
if plot_types:
context_parts.append(" + ".join(plot_types))
color_by = []
if "cell_type" in code_lower or "celltype" in code_lower or "tier3" in code_lower:
color_by.append("cell type")
if "leiden" in code_lower:
color_by.append("Leiden clusters")
if "louvain" in code_lower:
color_by.append("Leiden clusters")
if "batch" in code_lower or "sample" in code_lower:
color_by.append("batch/sample")
if "condition" in code_lower or "sample_type" in code_lower:
color_by.append("condition/disease stage")
if "leiden_neigh" in code_lower or "neighborhood" in code_lower or "neigh" in code_lower:
color_by.append("spatial neighborhood")
if "niche" in code_lower:
color_by.append("tissue niche")
if color_by:
context_parts.append(f"colored/grouped by: {', '.join(color_by)}")
title_patterns = [
r'plt\.title\s*\(\s*[\'"]([^\'"]+)[\'"]',
r'\.set_title\s*\(\s*[\'"]([^\'"]+)[\'"]',
r'title\s*=\s*[\'"]([^\'"]+)[\'"]',
]
for pattern in title_patterns:
match = re.search(pattern, code)
if match:
context_parts.append(f"Title: {match.group(1)}")
break
gene_patterns = [
r'var_names\s*=\s*\[([^\]]+)\]',
r'genes\s*=\s*\[([^\]]+)\]',
r"color\s*=\s*['\"]([A-Z][A-Z0-9]+)['\"]",
]
for pattern in gene_patterns:
match = re.search(pattern, code, re.IGNORECASE)
if match:
genes = match.group(1).strip()
if len(genes) < 200:
context_parts.append(f"Genes: {genes}")
break
comment_pattern = r'#\s*(.+?)$'
comments = re.findall(comment_pattern, code, re.MULTILINE)
relevant_comments = [c.strip() for c in comments if len(c.strip()) > 10 and len(c.strip()) < 100]
if relevant_comments:
context_parts.append(f"Code comments: {'; '.join(relevant_comments[:2])}")
if "comparison" in code_lower or "vs" in code_lower or "versus" in code_lower:
context_parts.append("Comparative analysis")
if "composition" in code_lower:
context_parts.append("Composition analysis")
if "proportion" in code_lower or "percentage" in code_lower:
context_parts.append("Proportion/percentage analysis")
if "dynamics" in code_lower or "trajectory" in code_lower:
context_parts.append("Dynamics/trajectory analysis")
if "interaction" in code_lower:
context_parts.append("Cell-cell interaction analysis")
if user_query:
query_truncated = user_query[:500] if len(user_query) > 500 else user_query
context_parts.append(f"Biological context: {query_truncated}")
return " | ".join(context_parts)
def run(self, user_query: str, config: Dict[str, Any] = None) -> Dict[str, Any]:
"""
Run the agent with a user query.
Args:
user_query: The user's task/question
config: Optional configuration dict
Returns:
Final agent state
"""
if config is None:
config = {"recursion_limit": 50}
elif "recursion_limit" not in config:
config["recursion_limit"] = 50
thread_id = config.get("thread_id", self.default_thread_id)
langgraph_config = {
"recursion_limit": config.get("recursion_limit", 50),
"configurable": {
"thread_id": thread_id
}
}
import sys
print(f"\033[1m<user query>\033[0m\n{user_query.strip()}\n\033[1m</user query>\033[0m\n")
sys.stdout.flush()
try:
existing_messages = self.conversation_history.get(thread_id, [])
if existing_messages:
print(f"✅ Restored {len(existing_messages)} messages from memory for thread: {thread_id}", flush=True)
else:
print(f"🔄 Starting new conversation thread: {thread_id}", flush=True)
initial_state = {
"messages": existing_messages + [HumanMessage(content=user_query)],
}
final_state = None
step_count = 0
all_messages_accumulated = []
for state_update in self.agent.stream(initial_state, stream_mode="values", config=langgraph_config):
step_count += 1
messages = state_update.get("messages", [])
all_messages_accumulated = messages
logging.debug(f"Stream step {step_count}: {len(messages)} messages")
for msg in messages[-1:]:
self._print_message(msg)
final_state = state_update
msg_content = messages[-1].content if messages and hasattr(messages[-1], 'content') else ""
if isinstance(msg_content, str) and "<conclude>" in msg_content:
break
if all_messages_accumulated:
self.conversation_history[thread_id] = all_messages_accumulated
logging.debug(f"Stream loop ended after {step_count} steps")
return {"messages": all_messages_accumulated}
except Exception as e:
print(f"Error: {e}", flush=True)
raise
def _print_message(self, message: BaseMessage):
"""Print a message with appropriate formatting."""
import sys
if isinstance(message, HumanMessage):
return
msg = message.content
if not msg or not isinstance(msg, str):
return
msg_stripped = msg.strip()
if not msg_stripped or msg_stripped == "[]" or msg_stripped == "{}":
return
if msg_stripped.startswith("[System]"):
return
def format_tag(text, tag, color_code):
pattern = rf"<{tag}>(.*?)</{tag}>"
def replacer(match):
content = match.group(1).strip()
return f"{color_code}<{tag}>\033[0m\n{content}\n{color_code}</{tag}>\033[0m"
return re.sub(pattern, replacer, text, flags=re.DOTALL)
conclude_match = re.search(r"<conclude>(.*?)</conclude>", msg, re.DOTALL)
if conclude_match:
try:
from rich.console import Console
from rich.markdown import Markdown
from rich.theme import Theme
custom_theme = Theme({
"markdown.code": "bold cyan",
"markdown.code_block": "cyan on grey93",
})
console = Console(theme=custom_theme, force_terminal=True)
conclude_content = conclude_match.group(1).strip()
if conclude_match.start() > 0:
pre_conclude = msg[:conclude_match.start()].strip()
pre_conclude = format_tag(pre_conclude, "act", "\033[91m")
pre_conclude = format_tag(pre_conclude, "observation", "\033[94m")
print(pre_conclude)
print()
sys.stdout.flush()
print("\033[1m<conclude>\033[0m")
md = Markdown(conclude_content, code_theme="github-light", inline_code_theme="cyan")
console.print(md)
print("\033[1m</conclude>\033[0m")
print()
sys.stdout.flush()
except ImportError:
display_msg = format_tag(msg_stripped, "act", "\033[91m")
display_msg = format_tag(display_msg, "observation", "\033[94m")
display_msg = format_tag(display_msg, "conclude", "\033[1m")
print(display_msg)
print()
sys.stdout.flush()
else:
display_msg = format_tag(msg_stripped, "act", "\033[91m")
display_msg = format_tag(display_msg, "observation", "\033[94m")
print(display_msg)
print()
sys.stdout.flush()