File size: 8,443 Bytes
6a55750
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
"""
Inference Script: RAG-Augmented Code Generation with Tool Calling

Combines:
  - Fine-tuned Qwen2.5-Coder-7B (code gen + tool-calling)
  - RAG pipeline (AST-aware code search)
  - ReAct-style reasoning loop for multi-step tasks

Usage:
  python inference.py --model your-username/qwen25-coder-7b-code-toolcall --repo /path/to/codebase
  python inference.py --model your-username/qwen25-coder-7b-code-toolcall --repo /path/to/codebase \
      --query "Implement a caching layer for the user service"
"""

import os, json, re, subprocess
from typing import Optional
from rag_pipeline import CodebaseIndexer, CodeRetriever, ContextBuilder

TOOLS = [
    {"type": "function", "function": {"name": "search_codebase",
        "description": "Search the internal codebase using semantic search.",
        "parameters": {"type": "object", "properties": {
            "query": {"type": "string"}, "file_pattern": {"type": "string"},
            "max_results": {"type": "integer", "default": 5}}, "required": ["query"]}}},
    {"type": "function", "function": {"name": "execute_python",
        "description": "Execute Python code in a sandboxed environment.",
        "parameters": {"type": "object", "properties": {
            "code": {"type": "string"}, "timeout": {"type": "integer", "default": 30}},
            "required": ["code"]}}},
    {"type": "function", "function": {"name": "read_file",
        "description": "Read contents of a file from the codebase.",
        "parameters": {"type": "object", "properties": {
            "path": {"type": "string"}, "start_line": {"type": "integer"},
            "end_line": {"type": "integer"}}, "required": ["path"]}}},
    {"type": "function", "function": {"name": "run_tests",
        "description": "Run unit tests for a module or file.",
        "parameters": {"type": "object", "properties": {
            "test_path": {"type": "string"}, "verbose": {"type": "boolean", "default": True}},
            "required": ["test_path"]}}}
]

SYSTEM_PROMPT = '''You are an expert Python programmer with access to our internal codebase via tools.
When helping with code tasks:
1. ALWAYS search the codebase first to understand existing patterns
2. Reference and integrate with existing code
3. Follow the codebase's conventions
4. Write comprehensive docstrings and type hints
5. Consider edge cases and error handling
6. For complex tasks, break them into steps and use tools iteratively
Think step by step.'''


class ToolExecutor:
    def __init__(self, retriever: CodeRetriever, repo_path: str):
        self.retriever = retriever
        self.repo_path = repo_path

    def execute(self, tool_name: str, arguments: dict) -> str:
        handlers = {"search_codebase": self._search, "execute_python": self._execute,
                    "read_file": self._read, "run_tests": self._test}
        return handlers.get(tool_name, lambda **kw: f"Unknown tool: {tool_name}")(**arguments)

    def _search(self, query, file_pattern=None, max_results=5):
        import fnmatch
        results = self.retriever.search(query, top_k=max_results)
        output = [{"file": c.file_path, "name": c.name, "type": c.chunk_type,
                   "score": round(s, 3), "signature": c.signature,
                   "content": c.content[:500]} for c, s in results
                  if not file_pattern or fnmatch.fnmatch(c.file_path, file_pattern)]
        return json.dumps({"results": output}, indent=2)

    def _execute(self, code, timeout=30):
        try:
            r = subprocess.run(["python", "-c", code], capture_output=True, text=True,
                             timeout=timeout, cwd=self.repo_path)
            return (r.stdout or "") + (r.stderr or "") or "Success (no output)"
        except subprocess.TimeoutExpired:
            return f"Timed out after {timeout}s"

    def _read(self, path, start_line=None, end_line=None):
        fp = os.path.join(self.repo_path, path)
        if not os.path.exists(fp): return f"File not found: {path}"
        with open(fp) as f: lines = f.readlines()
        if start_line: lines = lines[start_line-1:end_line]
        return "".join(lines)

    def _test(self, test_path, verbose=True):
        cmd = ["python", "-m", "pytest", test_path] + (["-v"] if verbose else [])
        try:
            r = subprocess.run(cmd, capture_output=True, text=True, timeout=120, cwd=self.repo_path)
            return r.stdout + r.stderr
        except subprocess.TimeoutExpired:
            return "Tests timed out"


class CodeAgent:
    """ReAct-style agent combining fine-tuned LLM with tool execution."""

    def __init__(self, model_id, retriever, repo_path, max_turns=10, device="auto"):
        self.tool_executor = ToolExecutor(retriever, repo_path)
        self.context_builder = ContextBuilder(retriever)
        self.max_turns = max_turns
        from transformers import AutoModelForCausalLM, AutoTokenizer
        self.tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
        self.model = AutoModelForCausalLM.from_pretrained(model_id, torch_dtype="auto",
            device_map=device, trust_remote_code=True)

    def run(self, user_query, current_file=None):
        rag_context = self.context_builder.build_context(user_query)
        user_content = user_query
        if rag_context:
            user_content += f"\\n\\n--- Code context ---\\n{rag_context}\\n--- End ---"
        messages = [{"role": "system", "content": SYSTEM_PROMPT},
                   {"role": "user", "content": user_content}]

        for turn in range(self.max_turns):
            response = self._generate(messages)
            tool_calls = self._extract_tool_calls(response)
            if not tool_calls:
                return response
            messages.append({"role": "assistant", "content": response})
            for tc in tool_calls:
                result = self.tool_executor.execute(tc["name"], tc["arguments"])
                messages.append({"role": "tool", "name": tc["name"], "content": result[:3000]})
        return response

    def _generate(self, messages):
        import torch
        text = self.tokenizer.apply_chat_template(messages, tools=TOOLS, tokenize=False,
                                                   add_generation_prompt=True)
        inputs = self.tokenizer(text, return_tensors="pt").to(self.model.device)
        with torch.no_grad():
            outputs = self.model.generate(**inputs, max_new_tokens=4096, temperature=0.7,
                                         top_p=0.9, do_sample=True)
        return self.tokenizer.decode(outputs[0][inputs["input_ids"].shape[-1]:],
                                    skip_special_tokens=True).strip()

    def _extract_tool_calls(self, response):
        tool_calls = []
        for pattern in [r'\\{"name":\\s*"(\\w+)",\\s*"arguments":\\s*(\\{[^}]+\\})\\}',
                       r'<tool_call>\\n?(.*?)\\n?</tool_call>']:
            for m in re.finditer(pattern, response, re.DOTALL):
                try:
                    if m.lastindex == 2:
                        tool_calls.append({"name": m.group(1), "arguments": json.loads(m.group(2))})
                    else:
                        data = json.loads(m.group(1))
                        if "name" in data: tool_calls.append(data)
                except: pass
        return tool_calls


if __name__ == "__main__":
    import argparse
    parser = argparse.ArgumentParser()
    parser.add_argument("--model", required=True)
    parser.add_argument("--repo", required=True)
    parser.add_argument("--query", "-q", default=None)
    parser.add_argument("--index-dir", default=None)
    parser.add_argument("--embedding-model", default="jinaai/jina-embeddings-v2-base-code")
    args = parser.parse_args()

    if args.index_dir and os.path.exists(os.path.join(args.index_dir, "chunks.json")):
        retriever = CodeRetriever(args.embedding_model)
        retriever.load_index(args.index_dir)
    else:
        retriever = CodebaseIndexer(args.repo, embedding_model=args.embedding_model).index()
        if args.index_dir: retriever.save_index(args.index_dir)

    agent = CodeAgent(args.model, retriever, args.repo)
    if args.query:
        print(agent.run(args.query))
    else:
        print("Code Assistant (type 'quit' to exit)")
        while True:
            q = input("> ").strip()
            if q.lower() in ("quit", "exit", "q"): break
            if q: print(agent.run(q))