From 052b7fe0d1989af50eb672aef6cb65ff19247462 Mon Sep 17 00:00:00 2001
From: Riccardo Schirone <562321+ret2libc@users.noreply.github.com>
Date: Tue, 3 Jun 2025 10:23:30 +0200
Subject: [PATCH] patcher: extract tools in separate files for reusability
(#765)
---
.../src/buttercup/patcher/agents/common.py | 10 +
.../patcher/agents/context_retriever.py | 395 ++----------------
patcher/src/buttercup/patcher/agents/tools.py | 264 ++++++++++++
patcher/tests/test_context_retriever.py | 32 +-
4 files changed, 302 insertions(+), 399 deletions(-)
create mode 100644 patcher/src/buttercup/patcher/agents/tools.py
diff --git a/patcher/src/buttercup/patcher/agents/common.py b/patcher/src/buttercup/patcher/agents/common.py
index 4f23e9b9..249b5e7a 100644
--- a/patcher/src/buttercup/patcher/agents/common.py
+++ b/patcher/src/buttercup/patcher/agents/common.py
@@ -3,6 +3,7 @@
from __future__ import annotations
from enum import Enum
+from pathlib import Path
from buttercup.common.llm import ButtercupLLM, create_default_llm_with_temperature
from typing import Annotated, Sequence
from buttercup.common.clusterfuzz_parser import CrashInfo
@@ -16,6 +17,7 @@ from langchain_core.output_parsers import StrOutputParser
from pydantic import BaseModel, Field
from buttercup.patcher.utils import PatchInput, PatchOutput, CHAIN_CALL_TYPE
from buttercup.common.challenge_task import ChallengeTask
+from langgraph.prebuilt.chat_agent_executor import AgentStatePydantic
import re
import uuid
@@ -163,6 +165,14 @@ class ExecutionInfo(BaseModel):
code_snippet_requests: list[CodeSnippetRequest] = Field(default_factory=list)
+class BaseCtxState(AgentStatePydantic):
+ """Base state for the context retriever agents"""
+
+ challenge_task_dir: Path
+ work_dir: Path
+ challenge_task_dir_ro: Path | None = None
+
+
class PatcherAgentState(BaseModel):
"""State for the Patcher Agent."""
diff --git a/patcher/src/buttercup/patcher/agents/context_retriever.py b/patcher/src/buttercup/patcher/agents/context_retriever.py
index dec8de93..e15ac98d 100644
--- a/patcher/src/buttercup/patcher/agents/context_retriever.py
+++ b/patcher/src/buttercup/patcher/agents/context_retriever.py
@@ -15,13 +15,11 @@ from langchain_core.output_parsers import StrOutputParser
from buttercup.common.stack_parsing import parse_stacktrace
from langgraph.checkpoint.memory import InMemorySaver
from langchain_core.messages import AnyMessage, SystemMessage, HumanMessage, AIMessage
-from functools import lru_cache
-from buttercup.common.challenge_task import CommandResult
from buttercup.patcher.agents.config import PatcherConfig
from dataclasses import dataclass, field
from langchain_core.messages import ToolMessage
from typing import Annotated, Any, Literal
-from pydantic import Field, ValidationError, BaseModel
+from pydantic import Field, ValidationError
from pathlib import Path
from langchain_core.tools import tool
from langgraph.prebuilt import InjectedState
@@ -30,7 +28,6 @@ from buttercup.common.challenge_task import ChallengeTask
from langchain_core.runnables import Runnable, RunnableConfig
from langchain_core.runnables.config import get_executor_for_config
-from buttercup.program_model.utils.common import Function, TypeDefinition
from buttercup.patcher.utils import truncate_output
from buttercup.patcher.agents.common import (
PatcherAgentBase,
@@ -40,17 +37,29 @@ from buttercup.patcher.agents.common import (
CodeSnippetRequest,
PatcherAgentName,
PatcherAgentState,
+ BaseCtxState,
)
from buttercup.common.llm import ButtercupLLM, create_default_llm
from langgraph.types import Command
-from langgraph.prebuilt.chat_agent_executor import AgentStatePydantic, create_react_agent
+from langgraph.prebuilt.chat_agent_executor import create_react_agent
+from buttercup.patcher.agents.tools import (
+ get_challenge,
+ ls,
+ grep,
+ cat,
+ get_lines,
+ MAX_OUTPUT_LENGTH,
+ get_function_tool_impl,
+ get_type_tool_impl,
+ get_callees,
+ get_callers,
+ get_function,
+ get_type,
+)
-from buttercup.program_model.codequery import CodeQueryPersistent
-
logger = logging.getLogger(__name__)
-MAX_OUTPUT_LENGTH = 10000
SYSTEM_TMPL = """You are an agent - please keep going until the user’s query is completely resolved, before ending your turn and yielding back to the user. Only terminate your turn when you are sure that the problem is solved.
If you are not sure about file content or codebase structure pertaining to the user’s request, use your tools to read files and gather the relevant information: do NOT guess or make up an answer.
You MUST plan extensively before each function call, and reflect extensively on the outcomes of the previous function calls. DO NOT do this entire process by making function calls only, as this can impair your ability to solve the problem and think insightfully.
@@ -107,32 +116,6 @@ Guidelines:
REMEMBER: The purpose of `track_snippet` is to **record already verified** code—not to discover or search.
"""
-CHECK_CODE_SNIPPETS_USER_MSG = """You are an AI assistant. Your job is to provide code snippets to the software engineer that answers the original request.
-
-The original request was:
-
-{REQUEST}
-
-
-Here is the list of code snippets you have found:
-
-{CODE_SNIPPETS}
-
-
-Did you find all the code snippets you need?
-Did you completely answer the original request?
-Do the code snippets fully answer the original request? (e.g. if the original \
-request was to find the code snippet that implements a function, make sure you \
-have found the code snippet that implements that function and that the function
-is fully implemented)
-"""
-
-CHECK_CODE_SNIPPETS_PROMPT = ChatPromptTemplate.from_messages(
- [
- ("user", CHECK_CODE_SNIPPETS_USER_MSG),
- ]
-)
-
DUPLICATE_CODE_SNIPPET_USER_MSG = """You are an AI assistant. Your job is to check if a code snippet request is already satisfied by the available code snippets.
Here is the code snippet request:
@@ -309,73 +292,19 @@ Validation:
"""
-class CheckCodeSnippetsOutput(BaseModel):
- """Output of the check_code_snippets tool"""
-
- reasoning: str = Field(
- description="Scratchpad space to reason about the code snippets you have found, the original request and whether you have found all the code snippets you need"
- )
- found_all: bool = Field(description="Whether you have found all the code snippets you need")
- fully_answered_request: bool = Field(description="Did you completely answer the original request?")
- fully_implemented: bool = Field(description="Do the code snippets fully answer the original request?")
- success: bool = Field(description="Can you consider the request satisfied?")
-
-
-class BaseCtxState(AgentStatePydantic):
- """Base state for the context retriever agents"""
-
- challenge_task_dir: Path
- challenge_task_dir_ro: Path | None = None
-
-
class CodeSnippetManagerState(BaseCtxState):
"""State for the code snippet manager"""
request: str
- work_dir: Path
code_snippets: Annotated[list[ContextCodeSnippet], operator.add] = Field(default_factory=list)
class FindTestsState(BaseCtxState):
"""State for the find tests agent."""
- work_dir: Path
tests_instructions: str | None = None
-@lru_cache(maxsize=100)
-def _get_challenge(task_dir: Path, task_dir_ro: Path | None = None) -> ChallengeTask:
- if task_dir_ro:
- return ChallengeTask(task_dir, local_task_dir=task_dir_ro)
-
- return ChallengeTask(task_dir, local_task_dir=task_dir)
-
-
-@lru_cache(maxsize=100)
-def _get_codequery(task_dir: Path, work_dir: Path) -> CodeQueryPersistent:
- challenge = _get_challenge(task_dir)
- return CodeQueryPersistent(challenge, work_dir=work_dir)
-
-
-def _wrap_command_output(command: str | list[str], cmd_res: CommandResult, output: str | None = None) -> str:
- if output is None:
- output = cmd_res.output.decode("utf-8")
-
- if isinstance(command, list):
- command = " ".join(command)
-
- return f"""
-{command}
-{cmd_res.returncode}
-
-{truncate_output(output, MAX_OUTPUT_LENGTH)}
-
-
-{truncate_output(cmd_res.error, MAX_OUTPUT_LENGTH)}
-
-"""
-
-
@tool
def think(reasoning: str) -> str:
"""Think more about the problem and the tools available to you."""
@@ -386,30 +315,6 @@ You should try to find the most relevant code snippet. Reasoning: \
with the newly discovered information?"""
-def _check_code_snippets(request: str, code_snippets: list[ContextCodeSnippet]) -> tuple[bool, str]:
- if len(code_snippets) == 0:
- return False, "No code snippets found. Please try again."
-
- llm = create_default_llm(model_name=ButtercupLLM.OPENAI_GPT_4_1_MINI.value)
- fallback_llms = [
- create_default_llm(model_name=ButtercupLLM.CLAUDE_3_5_SONNET.value),
- ]
- llm = llm.with_fallbacks(fallback_llms)
- check_code_snippets_chain = CHECK_CODE_SNIPPETS_PROMPT | llm.with_structured_output(CheckCodeSnippetsOutput)
- try:
- res: CheckCodeSnippetsOutput = check_code_snippets_chain.invoke(
- {
- "REQUEST": request,
- "CODE_SNIPPETS": "\n".join(str(code_snippet) for code_snippet in code_snippets),
- }
- )
- except Exception as e:
- logger.error("Error checking code snippets: %s", e)
- return False, "Error checking code snippets. Please try again."
-
- return res.success, res.reasoning
-
-
def _return_command_tool_message(
tool_call_id: str, message: str, new_code_snippets: list[ContextCodeSnippet] | None = None
) -> Command:
@@ -422,86 +327,6 @@ def _return_command_tool_message(
return Command(update=update_state)
-def _check_snippets_and_return(
- state: CodeSnippetManagerState, code_snippets: list[ContextCodeSnippet], tool_call_id: str
-) -> Command:
- assert code_snippets, "code_snippets must be non-empty"
- success, reasoning = _check_code_snippets(state.request, state.code_snippets + code_snippets)
- if not success:
- msg = f"Found {len(code_snippets)} code snippets, but the request has NOT been satisfied yet. You MUST keep going. \n\nReasoning: {reasoning}"
- return _return_command_tool_message(tool_call_id, msg, code_snippets)
-
- return _return_command_tool_message(
- tool_call_id,
- f"Found {len(code_snippets)} code snippets that satisfies the user's request. You MUST STOP NOW. You do not need to call any other tool.",
- code_snippets,
- )
-
-
-@tool
-def ls(
- file_path: str,
- *,
- state: Annotated[BaseCtxState, InjectedState],
-) -> str:
- """List the files in the given file_path in the project's source directory."""
- path = Path(file_path)
- logger.info("Listing files in %s", path)
- args = ["ls", "-la"]
- if path:
- args.append(str(path))
- challenge = _get_challenge(state.challenge_task_dir)
- ls_cmd_res = challenge.exec_docker_cmd(args)
- return _wrap_command_output(args, ls_cmd_res)
-
-
-@tool
-def grep(
- pattern: str,
- file_path: str | None,
- state: Annotated[BaseCtxState, InjectedState],
-) -> str:
- """Grep for a string and return a 5-line context around the match, together \
- with line numbers. If no file_path is provided, search the entire project. \
- Prefer using this tool over cat. If you need to search several files, just \
- call call this tool without any file_path."""
- path = Path(file_path) if file_path else None
- logger.info("Searching for %s in %s", pattern, path)
- args = ["grep", "-C", "5", "-nHrE", pattern]
- if path:
- args.append(str(path))
- challenge = _get_challenge(state.challenge_task_dir)
- grep_cmd_res = challenge.exec_docker_cmd(args)
- return _wrap_command_output(args, grep_cmd_res)
-
-
-@tool
-def cat(file_path: str, state: Annotated[BaseCtxState, InjectedState]) -> str:
- """Read the contents of a file. Use this tool only if grep and get_lines do not work as it might return a large amount of text."""
- path = Path(file_path)
- logger.info("Reading contents of %s", path)
- challenge = _get_challenge(state.challenge_task_dir)
- args = ["cat", str(path)]
- cat_cmd_res = challenge.exec_docker_cmd(args)
- return _wrap_command_output(args, cat_cmd_res)
-
-
-@tool
-def get_lines(
- file_path: str,
- start: int,
- end: int,
- state: Annotated[BaseCtxState, InjectedState],
-) -> str:
- """Get a range of lines from a file. Prefer using this tool over cat."""
- path = Path(file_path)
- logger.info("Getting lines %d-%d of %s", start, end, path)
- challenge = _get_challenge(state.challenge_task_dir)
- get_lines_res_cmd = challenge.exec_docker_cmd(["cat", str(path)])
- get_lines_output = get_lines_res_cmd.output.decode("utf-8").splitlines()[start:end]
- return _wrap_command_output(f"get_lines {path} {start} {end}", get_lines_res_cmd, "\n".join(get_lines_output))
-
-
@tool
def track_snippet(
file_path: str,
@@ -568,12 +393,12 @@ def track_snippet(
path = Path(file_path)
code_snippets = None
if function_name:
- code_snippets = _get_function(function_name, file_path, state)
+ code_snippets = get_function_tool_impl(function_name, file_path, state)
elif type_name:
- code_snippets = _get_type(type_name, file_path, state)
+ code_snippets = get_type_tool_impl(type_name, file_path, state)
elif start_line and end_line:
logger.info("Getting lines %d-%d of %s", start_line, end_line, path)
- challenge = _get_challenge(state.challenge_task_dir)
+ challenge = get_challenge(state.challenge_task_dir)
get_lines_res_cmd = challenge.exec_docker_cmd(["cat", str(path)])
# Get a few lines before and after the requested lines in case the LLM does
# small mistakes when requesting the lines
@@ -606,178 +431,6 @@ def track_snippet(
)
-def _get_codequery_function(codequery: CodeQueryPersistent, name: str, path: Path | None) -> Function:
- functions = codequery.get_functions(name, path)
- if not functions:
- raise ValueError(f"No definition found for function {name} in {path}")
-
- return functions[0]
-
-
-def _add_functions_code_snippets(functions: list[Function], suffix: str = "") -> list[ContextCodeSnippet]:
- return [
- ContextCodeSnippet(
- key=CodeSnippetKey(
- file_path=function.file_path.as_posix(),
- ),
- start_line=body.start_line,
- end_line=body.end_line,
- code=body.body,
- description=f"Implementation of function {function.name}{suffix}",
- )
- for function in functions
- for body in function.bodies
- ]
-
-
-def _add_type_definitions_code_snippets(type_definitions: list[TypeDefinition]) -> list[ContextCodeSnippet]:
- return [
- ContextCodeSnippet(
- key=CodeSnippetKey(
- file_path=type_def.file_path.as_posix(),
- ),
- code=type_def.definition,
- start_line=type_def.definition_line,
- end_line=type_def.definition_line + len(type_def.definition.splitlines()),
- description=f"Definition of type {type_def.name}",
- )
- for type_def in type_definitions
- ]
-
-
-def _clean_function_name(function_name: str) -> str:
- if function_name.startswith("OSS_FUZZ_"):
- return function_name[len("OSS_FUZZ_") :]
- return function_name
-
-
-def _get_function(
- function_name: str,
- file_path: str | None,
- state: Annotated[CodeSnippetManagerState, InjectedState],
-) -> list[ContextCodeSnippet]:
- path = Path(file_path) if file_path else None
- challenge = _get_challenge(state.challenge_task_dir)
- if path and not path.is_absolute():
- # If the path is not absolute, it is relative to the container workdir
- path = challenge.workdir_from_dockerfile().joinpath(path)
-
- logger.info("Getting function definition of %s in %s", function_name, path)
- function_name = _clean_function_name(function_name)
- codequery = _get_codequery(state.challenge_task_dir, state.work_dir)
- functions = codequery.get_functions(function_name, path)
- if not functions:
- functions = codequery.get_functions(function_name, path, fuzzy=True)
- if not functions:
- functions = codequery.get_functions(function_name, None)
- if not functions:
- functions = codequery.get_functions(function_name, None, fuzzy=True)
- if not functions:
- raise ValueError(f"No definition found for function {function_name} in {path}")
-
- return _add_functions_code_snippets(functions)
-
-
-@tool
-def get_function(
- function_name: str,
- file_path: str | None,
- state: Annotated[CodeSnippetManagerState, InjectedState],
-) -> str:
- """Get a function's definition. If available, pass a file_path, \
- otherwise pass None. Use this when you want to get information about a \
- function. If not sure about the file path, pass None. Prefer using this \
- tool over any other and rely on others only if this tool fails or does \
- not work."""
- code_snippets = _get_function(function_name, file_path, state)
- output_str = "\n".join(str(code_snippet) for code_snippet in code_snippets)
- return output_str
-
-
-@tool
-def get_callers(
- function_name: str,
- file_path: str | None,
- state: Annotated[CodeSnippetManagerState, InjectedState],
-) -> str:
- """Get the callers of a function."""
- path = Path(file_path) if file_path else None
- logger.info("Getting callers of %s in %s", function_name, path)
- codequery = _get_codequery(state.challenge_task_dir, state.work_dir)
- function = _get_codequery_function(codequery, function_name, path)
- callers = codequery.get_callers(function)
- if not callers:
- raise ValueError(f"No callers found for function {function_name} in {path}")
-
- msg = f"""Found {len(callers)} callers of function {function_name}:
-{"\n".join(f"- `{caller.name}` in `{caller.file_path}`" for caller in callers)}
-
-If you need to get the definition of the caller in order to satisfy the request, \
-call `get_function` tool with the caller's name and file_path.
-"""
-
- return msg
-
-
-@tool
-def get_callees(
- function_name: str,
- file_path: str | None,
- state: Annotated[CodeSnippetManagerState, InjectedState],
-) -> str:
- """Get the callees of a function."""
- path = Path(file_path) if file_path else None
- logger.info("Getting callees of %s in %s", function_name, path)
- codequery = _get_codequery(state.challenge_task_dir, state.work_dir)
- function = _get_codequery_function(codequery, function_name, path)
- callees = codequery.get_callees(function)
- if not callees:
- raise ValueError(f"No callees found for function {function_name} in {path}")
-
- msg = f"""Found {len(callees)} callees of function {function_name}:
-{"\n".join(f"- `{callee.name}` in `{callee.file_path}`" for callee in callees)}
-
-If you need to get the definition of the callee in order to satisfy the request, \
-call `get_function` tool with the callee's name and file_path.
-"""
-
- return msg
-
-
-def _get_type(
- type_name: str,
- file_path: str | None,
- state: Annotated[CodeSnippetManagerState, InjectedState],
-) -> list[ContextCodeSnippet]:
- path = Path(file_path) if file_path else None
-
- logger.info("Getting type definition of %s in %s", type_name, path)
- codequery = _get_codequery(state.challenge_task_dir, state.work_dir)
- types = codequery.get_types(type_name, path)
- if not types:
- types = codequery.get_types(type_name, path, fuzzy=True)
- if not types:
- raise ValueError(f"No definition found for type {type_name} in {path}")
-
- return _add_type_definitions_code_snippets(types)
-
-
-@tool
-def get_type(
- type_name: str,
- file_path: str | None,
- *,
- state: Annotated[CodeSnippetManagerState, InjectedState],
-) -> str:
- """Get a type/class/typedef/struct/enum/macro's definition. If available, pass a file_path, \
- otherwise pass None. Use this when you want to get information about a type. \
- If not sure about the file path, pass None. Prefer using this tool over any \
- other and rely on others only if this tool fails or does not work."""
- code_snippets = _get_type(type_name, file_path, state)
- output_str = "\n".join(str(code_snippet) for code_snippet in code_snippets)
- return output_str
-
-
@tool
def sh(
command: str,
@@ -805,7 +458,7 @@ def sh(
command_file_path = Path(f.name)
command_file_path.chmod(0o755)
- challenge = _get_challenge(state.challenge_task_dir, state.challenge_task_dir_ro)
+ challenge = get_challenge(state.challenge_task_dir, state.challenge_task_dir_ro)
sh_cmd_res = challenge.exec_docker_cmd(
"/tmp/command.sh",
mount_dirs={
@@ -1038,7 +691,7 @@ class ContextRetrieverAgent(PatcherAgentBase):
self.ls_cwd = "ls cwd failed"
def _prompt(self, state: CodeSnippetManagerState) -> list[AnyMessage]:
- challenge = _get_challenge(state.challenge_task_dir)
+ challenge = get_challenge(state.challenge_task_dir)
return [
SystemMessage(content=SYSTEM_TMPL),
HumanMessage(
@@ -1065,7 +718,7 @@ class ContextRetrieverAgent(PatcherAgentBase):
]
def _find_tests_prompt(self, state: FindTestsState) -> list[AnyMessage]:
- challenge = _get_challenge(state.challenge_task_dir, state.challenge_task_dir_ro)
+ challenge = get_challenge(state.challenge_task_dir, state.challenge_task_dir_ro)
ls_cwd = challenge.exec_docker_cmd(["ls", "-la"])
if ls_cwd.success:
@@ -1261,7 +914,7 @@ class ContextRetrieverAgent(PatcherAgentBase):
"""Get the initial context for the diff analysis."""
configuration = PatcherConfig.from_configurable(config)
stacktrace = parse_stacktrace(state.context.sanitizer_output)
- challenge = _get_challenge(state.context.challenge_task_dir)
+ challenge = get_challenge(state.context.challenge_task_dir)
# Request code snippet for the first two functions in the stacktrace
logger.info("[%s] Getting initial context from stacktrace", self.challenge.task_meta.task_id)
diff --git a/patcher/src/buttercup/patcher/agents/tools.py b/patcher/src/buttercup/patcher/agents/tools.py
new file mode 100644
index 00000000..bce356a0
--- /dev/null
+++ b/patcher/src/buttercup/patcher/agents/tools.py
@@ -0,0 +1,264 @@
+"""Tools for the patcher agents."""
+
+from __future__ import annotations
+
+import logging
+from functools import lru_cache
+from buttercup.common.challenge_task import CommandResult
+from buttercup.program_model.codequery import CodeQueryPersistent
+from buttercup.program_model.utils.common import Function, TypeDefinition
+from typing import Annotated
+from pathlib import Path
+from langchain_core.tools import tool
+from langgraph.prebuilt import InjectedState
+from buttercup.common.challenge_task import ChallengeTask
+from buttercup.patcher.utils import truncate_output
+from buttercup.patcher.agents.common import BaseCtxState, ContextCodeSnippet, CodeSnippetKey
+
+logger = logging.getLogger(__name__)
+
+MAX_OUTPUT_LENGTH = 10000
+
+
+@lru_cache(maxsize=100)
+def get_challenge(task_dir: Path, task_dir_ro: Path | None = None) -> ChallengeTask:
+ if task_dir_ro:
+ return ChallengeTask(task_dir, local_task_dir=task_dir_ro)
+
+ return ChallengeTask(task_dir, local_task_dir=task_dir)
+
+
+@lru_cache(maxsize=100)
+def get_codequery(task_dir: Path, work_dir: Path) -> CodeQueryPersistent:
+ challenge = get_challenge(task_dir)
+ return CodeQueryPersistent(challenge, work_dir=work_dir)
+
+
+def _wrap_command_output(command: str | list[str], cmd_res: CommandResult, output: str | None = None) -> str:
+ if output is None:
+ output = cmd_res.output.decode("utf-8")
+
+ if isinstance(command, list):
+ command = " ".join(command)
+
+ return f"""
+{command}
+{cmd_res.returncode}
+
+{truncate_output(output, MAX_OUTPUT_LENGTH)}
+
+
+{truncate_output(cmd_res.error, MAX_OUTPUT_LENGTH)}
+
+"""
+
+
+@tool
+def ls(
+ file_path: str,
+ *,
+ state: Annotated[BaseCtxState, InjectedState],
+) -> str:
+ """List the files in the given file_path in the project's source directory."""
+ path = Path(file_path)
+ logger.info("Listing files in %s", path)
+ args = ["ls", "-la"]
+ if path:
+ args.append(str(path))
+ challenge = get_challenge(state.challenge_task_dir)
+ ls_cmd_res = challenge.exec_docker_cmd(args)
+ return _wrap_command_output(args, ls_cmd_res)
+
+
+@tool
+def grep(
+ pattern: str,
+ file_path: str | None,
+ state: Annotated[BaseCtxState, InjectedState],
+) -> str:
+ """Grep for a string and return a 5-line context around the match, together \
+ with line numbers. If no file_path is provided, search the entire project. \
+ Prefer using this tool over cat. If you need to search several files, just \
+ call call this tool without any file_path."""
+ path = Path(file_path) if file_path else None
+ logger.info("Searching for %s in %s", pattern, path)
+ args = ["grep", "-C", "5", "-nHrE", pattern]
+ if path:
+ args.append(str(path))
+ challenge = get_challenge(state.challenge_task_dir)
+ grep_cmd_res = challenge.exec_docker_cmd(args)
+ return _wrap_command_output(args, grep_cmd_res)
+
+
+@tool
+def cat(file_path: str, state: Annotated[BaseCtxState, InjectedState]) -> str:
+ """Read the contents of a file. Use this tool only if grep and get_lines do not work as it might return a large amount of text."""
+ path = Path(file_path)
+ logger.info("Reading contents of %s", path)
+ challenge = get_challenge(state.challenge_task_dir)
+ args = ["cat", str(path)]
+ cat_cmd_res = challenge.exec_docker_cmd(args)
+ return _wrap_command_output(args, cat_cmd_res)
+
+
+@tool
+def get_lines(
+ file_path: str,
+ start: int,
+ end: int,
+ state: Annotated[BaseCtxState, InjectedState],
+) -> str:
+ """Get a range of lines from a file. Prefer using this tool over cat."""
+ path = Path(file_path)
+ logger.info("Getting lines %d-%d of %s", start, end, path)
+ challenge = get_challenge(state.challenge_task_dir)
+ get_lines_res_cmd = challenge.exec_docker_cmd(["cat", str(path)])
+ get_lines_output = get_lines_res_cmd.output.decode("utf-8").splitlines()[start:end]
+ return _wrap_command_output(f"get_lines {path} {start} {end}", get_lines_res_cmd, "\n".join(get_lines_output))
+
+
+def _get_codequery_function(codequery: CodeQueryPersistent, name: str, path: Path | None) -> Function:
+ functions = codequery.get_functions(name, path)
+ if not functions:
+ raise ValueError(f"No definition found for function {name} in {path}")
+
+ return functions[0]
+
+
+def _add_functions_code_snippets(functions: list[Function], suffix: str = "") -> list[ContextCodeSnippet]:
+ return [
+ ContextCodeSnippet(
+ key=CodeSnippetKey(
+ file_path=function.file_path.as_posix(),
+ ),
+ start_line=body.start_line,
+ end_line=body.end_line,
+ code=body.body,
+ description=f"Implementation of function {function.name}{suffix}",
+ )
+ for function in functions
+ for body in function.bodies
+ ]
+
+
+def _add_type_definitions_code_snippets(type_definitions: list[TypeDefinition]) -> list[ContextCodeSnippet]:
+ return [
+ ContextCodeSnippet(
+ key=CodeSnippetKey(
+ file_path=type_def.file_path.as_posix(),
+ ),
+ code=type_def.definition,
+ start_line=type_def.definition_line,
+ end_line=type_def.definition_line + len(type_def.definition.splitlines()),
+ description=f"Definition of type {type_def.name}",
+ )
+ for type_def in type_definitions
+ ]
+
+
+def _clean_function_name(function_name: str) -> str:
+ if function_name.startswith("OSS_FUZZ_"):
+ return function_name[len("OSS_FUZZ_") :]
+ return function_name
+
+
+def get_function_tool_impl(function_name: str, file_path: str | None, state: BaseCtxState) -> list[ContextCodeSnippet]:
+ path = Path(file_path) if file_path else None
+ challenge = get_challenge(state.challenge_task_dir)
+ if path and not path.is_absolute():
+ # If the path is not absolute, it is relative to the container workdir
+ path = challenge.workdir_from_dockerfile().joinpath(path)
+
+ logger.info("Getting function definition of %s in %s", function_name, path)
+ function_name = _clean_function_name(function_name)
+ codequery = get_codequery(state.challenge_task_dir, state.work_dir)
+ functions = codequery.get_functions(function_name, path)
+ if not functions:
+ functions = codequery.get_functions(function_name, path, fuzzy=True)
+ if not functions:
+ functions = codequery.get_functions(function_name, None)
+ if not functions:
+ functions = codequery.get_functions(function_name, None, fuzzy=True)
+ if not functions:
+ raise ValueError(f"No definition found for function {function_name} in {path}")
+
+ return _add_functions_code_snippets(functions)
+
+
+@tool
+def get_function(function_name: str, file_path: str | None, *, state: Annotated[BaseCtxState, InjectedState]) -> str:
+ """Get a function's definition. If available, pass a file_path, \
+ otherwise pass None. Use this when you want to get information about a \
+ function. If not sure about the file path, pass None. Prefer using this \
+ tool over any other and rely on others only if this tool fails or does \
+ not work."""
+ code_snippets = get_function_tool_impl(function_name, file_path, state)
+ output_str = "\n".join(str(code_snippet) for code_snippet in code_snippets)
+ return output_str
+
+
+@tool
+def get_callers(function_name: str, file_path: str | None, *, state: Annotated[BaseCtxState, InjectedState]) -> str:
+ """Get the callers of a function."""
+ path = Path(file_path) if file_path else None
+ logger.info("Getting callers of %s in %s", function_name, path)
+ codequery = get_codequery(state.challenge_task_dir, state.work_dir)
+ function = _get_codequery_function(codequery, function_name, path)
+ callers = codequery.get_callers(function)
+ if not callers:
+ raise ValueError(f"No callers found for function {function_name} in {path}")
+
+ msg = f"""Found {len(callers)} callers of function {function_name}:
+{"\n".join(f"- `{caller.name}` in `{caller.file_path}`" for caller in callers)}
+
+If you need to get the definition of the caller in order to satisfy the request, \
+call `get_function` tool with the caller's name and file_path.
+"""
+
+ return msg
+
+
+@tool
+def get_callees(function_name: str, file_path: str | None, *, state: Annotated[BaseCtxState, InjectedState]) -> str:
+ """Get the callees of a function."""
+ path = Path(file_path) if file_path else None
+ logger.info("Getting callees of %s in %s", function_name, path)
+ codequery = get_codequery(state.challenge_task_dir, state.work_dir)
+ function = _get_codequery_function(codequery, function_name, path)
+ callees = codequery.get_callees(function)
+ if not callees:
+ raise ValueError(f"No callees found for function {function_name} in {path}")
+
+ msg = f"""Found {len(callees)} callees of function {function_name}:
+{"\n".join(f"- `{callee.name}` in `{callee.file_path}`" for callee in callees)}
+
+If you need to get the definition of the callee in order to satisfy the request, \
+call `get_function` tool with the callee's name and file_path.
+"""
+
+ return msg
+
+
+def get_type_tool_impl(type_name: str, file_path: str | None, state: BaseCtxState) -> list[ContextCodeSnippet]:
+ path = Path(file_path) if file_path else None
+
+ logger.info("Getting type definition of %s in %s", type_name, path)
+ codequery = get_codequery(state.challenge_task_dir, state.work_dir)
+ types = codequery.get_types(type_name, path)
+ if not types:
+ types = codequery.get_types(type_name, path, fuzzy=True)
+ if not types:
+ raise ValueError(f"No definition found for type {type_name} in {path}")
+
+ return _add_type_definitions_code_snippets(types)
+
+
+@tool
+def get_type(type_name: str, file_path: str | None, *, state: Annotated[BaseCtxState, InjectedState]) -> str:
+ """Get a type/class/typedef/struct/enum/macro's definition. If available, pass a file_path, \
+ otherwise pass None. Use this when you want to get information about a type. \
+ If not sure about the file path, pass None. Prefer using this tool over any \
+ other and rely on others only if this tool fails or does not work."""
+ code_snippets = get_type_tool_impl(type_name, file_path, state)
+ output_str = "\n".join(str(code_snippet) for code_snippet in code_snippets)
+ return output_str
diff --git a/patcher/tests/test_context_retriever.py b/patcher/tests/test_context_retriever.py
index 2ebc9df8..595e924e 100644
--- a/patcher/tests/test_context_retriever.py
+++ b/patcher/tests/test_context_retriever.py
@@ -15,7 +15,6 @@ from typing import Iterator
from langchain_core.language_models import BaseChatModel
from buttercup.patcher.agents.context_retriever import (
ContextRetrieverAgent,
- CheckCodeSnippetsOutput,
)
from buttercup.patcher.agents.common import (
@@ -94,32 +93,10 @@ def mock_duplicate_code_snippet_prompt(mock_cheap_llm: MagicMock):
return prompt
-@pytest.fixture
-def mock_check_code_snippet_llm():
- prompt = MagicMock(spec=Runnable)
- llm = MagicMock(spec=BaseChatModel)
- llm.with_fallbacks.return_value = llm
- llm.configurable_fields.return_value = llm
- llm.with_structured_output.return_value = llm
-
- llm.invoke.return_value = CheckCodeSnippetsOutput(
- reasoning="",
- found_all=True,
- fully_answered_request=True,
- fully_implemented=True,
- success=True,
- )
-
- prompt.__or__.side_effect = lambda other: llm
- llm.__or__.side_effect = lambda other: llm
- return prompt
-
-
@pytest.fixture(autouse=True)
def mock_llm_functions(
mock_agent_llm: MagicMock,
mock_cheap_llm: MagicMock,
- mock_check_code_snippet_llm: MagicMock,
mock_duplicate_code_snippet_prompt: MagicMock,
):
"""Mock LLM creation functions and environment variables."""
@@ -132,7 +109,6 @@ def mock_llm_functions(
import buttercup.patcher.agents.context_retriever
buttercup.patcher.agents.context_retriever.DUPLICATE_CODE_SNIPPET_PROMPT = mock_duplicate_code_snippet_prompt
- buttercup.patcher.agents.context_retriever.CHECK_CODE_SNIPPETS_PROMPT = mock_check_code_snippet_llm
yield
@@ -1358,7 +1334,7 @@ def test_llm_error_recovery(mock_agent: ContextRetrieverAgent, mock_agent_llm: M
assert code_snippet.code == "struct ebitmap_t { int a; }"
-@patch("buttercup.patcher.agents.context_retriever._get_challenge")
+@patch("buttercup.patcher.agents.context_retriever.get_challenge")
def test_get_initial_context_filters_llvm_frames(
mock_get_challenge: MagicMock,
mock_agent: ContextRetrieverAgent,
@@ -1428,7 +1404,7 @@ def test_get_initial_context_filters_llvm_frames(
)
-@patch("buttercup.patcher.agents.context_retriever._get_challenge")
+@patch("buttercup.patcher.agents.context_retriever.get_challenge")
def test_get_initial_context_includes_llvm_frames(
mock_get_challenge: MagicMock,
mock_agent: ContextRetrieverAgent,
@@ -1512,7 +1488,7 @@ def test_get_initial_context_includes_llvm_frames(
)
-@patch("buttercup.patcher.agents.context_retriever._get_challenge")
+@patch("buttercup.patcher.agents.context_retriever.get_challenge")
def test_get_initial_context_respects_n_initial_stackframes(
mock_get_challenge: MagicMock,
mock_agent: ContextRetrieverAgent,
@@ -1608,7 +1584,7 @@ def test_get_initial_context_respects_n_initial_stackframes(
assert mock_agent.process_request.call_count == 2
-@patch("buttercup.patcher.agents.context_retriever._get_challenge")
+@patch("buttercup.patcher.agents.context_retriever.get_challenge")
def test_get_initial_context_handles_multiple_stackframes(
mock_get_challenge: MagicMock,
mock_agent: ContextRetrieverAgent,