Team Ai
Apppublic

VirtualLab/AutonomousResearcher

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
main.py347 linesDownload Raw Back to root
1import asyncio2import logging3import json4from contextlib import asynccontextmanager5from pathlib import Path6from fastapi import FastAPI, WebSocket, Query7from fastapi.responses import FileResponse, HTMLResponse8from fastapi.staticfiles import StaticFiles9import requests10import subprocess11import uvicorn12import os13from datetime import datetime14import time15try:16    import markdown17except ImportError:18    markdown = None  # if markdown is not installed, .md files will be shown as plain text19 20# Configuration21CONFIG = {22    "system": "huggingface",  # "local" or "huggingface"23    "url": "https://openrouter.ai/api/v1/chat/completions",24    "check_interval": 240,25    "max_models": 1,26    "interaction_log": "logs/interaction.log",27    "key_number": 0,28    "max_keys": 329}30 31# (Existing key configuration code remains unchanged)32 33if CONFIG["system"] == "huggingface":34    keys = []35    keys.append({"key_num": 1, "key": os.getenv("OPENROUTER")})36    keys.append({"key_num": 2, "key": os.getenv("OPENROUTER_FALLBACK")})37    keys.append({"key_num": 3, "key": os.getenv("OPENROUTER_FALLBACK_2")})38    openrouter_key = keys[CONFIG["key_number"]]39elif CONFIG["system"] == "local":40    from keys import OPENROUTER_KEYS41    openrouter_key = OPENROUTER_KEYS[CONFIG["key_number"]]42 43# Initialize logging44logging.basicConfig(45    level=logging.INFO,46    format="%(asctime)s - %(levelname)s - %(message)s",47)48 49# Define the folder that will hold the agent's files50AGENT_DIR = Path("agent_files")51AGENT_DIR.mkdir(exist_ok=True)52 53# --- ResearchAgent class and its methods remain the same ---54# (e.g. _initialize_system, _prepare_logging, query_model, etc.)55 56# After CONFIG is defined, add a global timestamp for next call57next_call_time = time.time() + CONFIG["check_interval"]58 59class ResearchAgent:60    def __init__(self):61        self.conversation = []62        self.task_finished = False63        self.selected_model = 064        self.log_update_event = asyncio.Event()65        self._initialize_system()66        self._prepare_logging()67 68    def _initialize_system(self):69        self.system_prompt = """You are an AI Agent expert on computational biology. 70You are capable of using tools for interacting with a Linux shell."""71        self.user_prompt = """Your goal is to discover real, verifiable data in the field of nerve regeneration. Do not generate simulated or hypothetical research. Persist in obtaining accurate, real-world data even if some tools fail. You are aware that, through the CLI, you can install tools, navigate the internet, and create and execute your own scripts and tools as needed. Continue your research relentlessly and use all available resources to achieve authentic results."""72        self.tools = [{73            "type": "function",74            "function": {75                "name": "execute_command",76                "description": "Execute a terminal command",77                "parameters": {78                    "type": "object",79                    "properties": {80                        "command": {"type": "string", "description": "Command to execute"}81                    },82                    "required": ["command"]83                }84            }85        }, {86            "type": "function",87            "function": {88                "name": "finish_task",89                "description": "Mark task completion",90                "parameters": {91                    "type": "object",92                    "properties": {93                        "completed": {"type": "boolean", "description": "Task completion status"}94                    },95                    "required": ["completed"]96                }97            }98        }]99        self.conversation.append({"role": "system", "content": self.system_prompt})100        self.conversation.append({"role": "user", "content": self.user_prompt})101        self.models = [102             {"name": "Gemini2Pro", "address": "google/gemini-2.0-pro-exp-02-05:free"},103             {"name": "Phi3", "address": "microsoft/phi-3-mini-128k-instruct:free"},104             {"name": "Gemini2Flash", "address": "google/gemini-2.0-flash-lite-preview-02-05:free"}105        ]106 107    def _prepare_logging(self):108        self.log_file = Path(CONFIG["interaction_log"])109        self.log_file.parent.mkdir(exist_ok=True)110        self.log_file.touch(exist_ok=True)111        self.log_file.write_text("")112 113    async def query_model(self):114        global next_call_time115        while not self.task_finished:116            try:117                response = await self._get_model_response()118                await self._process_response(response)119            except Exception as e:120                logging.error(f"Query failed: {e}")121            next_call_time = time.time() + CONFIG["check_interval"]122            await asyncio.sleep(CONFIG["check_interval"])123 124    async def _get_model_response(self):125        global keys, openrouter_key126        for attempt in range(CONFIG["max_models"]):127            logging.info(f"""POST request:128                         url: {CONFIG['url']},129                         key: {openrouter_key['key_num']},130                         model: {self.models[attempt]['name']}131                         """)132            try:133                response = requests.post(134                    url=CONFIG["url"],135                    headers={"Authorization": f"Bearer {openrouter_key['key']}"},136                    json={137                        "model": self.models[attempt]["address"],138                        "messages": self.conversation,139                        "tools": self.tools,140                        "tool_choice": "auto",141                        "temperature": 0.2142                    }143                )144                json_response = response.json()145                if "error" in json_response:146                    if json_response["error"]["code"] == 429:147                        logging.error("Quota exceeded")148                        continue149                    elif json_response["error"]["code"] == 400:150                        logging.error("Tool response not immediately after tool call")151                        logging.info(f"For Debug. Conversation: {self.conversation}")152                        return None153                    else:154                        logging.error(f"Server returned error: {json_response}")155                        logging.info(f"For Debug. Conversation: {self.conversation}")156                openrouter_key = keys[(CONFIG["key_number"] + 1) % CONFIG["max_keys"]]157                return json_response158            except Exception as e:159                logging.error(f"Error making the post request. {e}")160        return None161 162    async def _process_response(self, response):163        if not response:164            logging.error("No response")165        elif "choices" not in response:166            logging.error("No 'choices' in response")167            logging.info(f"Response {response}")168            return169        message = response["choices"][0]["message"]170        self.conversation.append(message)171        self._log_interaction(f"AI_TEXT: {message['content']}")172        if "tool_calls" in message:173            await self._handle_tool_calls(message["tool_calls"])174        else:175            self.conversation.append({"role": "user", "content": "Go"})176 177    async def _handle_tool_calls(self, tool_calls):178        for call in tool_calls:179            try:180                if call["function"]["name"] == "execute_command":181                    await self._handle_command(call)182                elif call["function"]["name"] == "finish_task":183                    self._handle_completion(call)184            except Exception as e:185                logging.error(f"Tool call failed: {e}")186                self._log_interaction(f"ERROR: {str(e)}")187 188    async def _handle_command(self, tool_call):189        args = json.loads(tool_call["function"]["arguments"])190        command = args.get("command", "")191        self._log_interaction(f"COMMAND: {command}")192        loop = asyncio.get_running_loop()193        output = await loop.run_in_executor(None, self._execute_command, command)194        self.conversation.append({195            "role": "tool",196            "tool_call_id": tool_call["id"],197            "name": "execute_command",198            "content": output199        })200        self._log_interaction(f"TERMINAL: {output}")201 202    def _handle_completion(self, tool_call):203        args = json.loads(tool_call["function"]["arguments"])204        if args.get("completed", False):205            self.task_finished = True206            logging.info("Research task marked as completed")207            self._log_interaction("Research task marked as completed")208 209    def _execute_command(self, command):210        try:211            result = subprocess.run(212                command,213                shell=True,214                text=True,215                capture_output=True,216                timeout=30217            )218            return result.stdout or result.stderr or "[No output]"219        except Exception as e:220            return f"Command failed: {str(e)}"221 222    def _log_interaction(self, content):223        entry_separator = "Log entry. Timestamp#:"224        timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S")225        content_separator = " Content#:"226        log_entry = f"{entry_separator}{timestamp}{content_separator}{content}"227        with self.log_file.open("a") as f:228            f.write(f"{log_entry}\n\n")229            f.flush()230        self.log_update_event.set()231 232@asynccontextmanager233async def lifespan(app: FastAPI):234    agent = ResearchAgent()235    app.state.agent = agent236    asyncio.create_task(agent.query_model())237    yield238    logging.info("Application shutdown")239 240app = FastAPI(lifespan=lifespan)241app.mount("/static", StaticFiles(directory="app"), name="static")242 243@app.get("/")244async def get_home():245    try:246        return HTMLResponse(Path("app/home.html").read_text())247    except Exception as e:248        return HTMLResponse("<h1>Research Interface</h1>")249 250@app.get("/log")251async def get_log():252    return FileResponse(CONFIG["interaction_log"])253 254# Return only the agent's files (in AGENT_DIR)255@app.get("/directory")256async def get_directory():257    def build_tree(path: Path, base: Path):258        tree = {"name": path.name, "relative_path": str(path.relative_to(base)), "type": "directory", "children": []}259        try:260            for child in sorted(path.iterdir(), key=lambda p: (p.is_file(), p.name.lower())):261                if child.is_dir():262                    tree["children"].append(build_tree(child, base))263                else:264                    tree["children"].append({265                        "name": child.name,266                        "relative_path": str(child.relative_to(base)),267                        "type": "file"268                    })269        except Exception as e:270            logging.error(f"Error reading directory {path}: {e}")271        return tree272    tree = build_tree(AGENT_DIR, AGENT_DIR)273    return tree274 275# Serve file content. If a markdown file, render it as HTML if markdown package is available.276@app.get("/file")277async def get_file(path: str = Query(...)):278    base_dir = AGENT_DIR279    file_path = base_dir / Path(path)280    # Prevent path traversal281    try:282        file_path.relative_to(base_dir)283    except ValueError:284        return HTMLResponse("Access denied", status_code=403)285    if file_path.exists() and file_path.is_file():286        try:287            content = file_path.read_text()288            if file_path.suffix.lower() == ".md" and markdown:289                html_content = markdown.markdown(content)290                return HTMLResponse(html_content)291            return HTMLResponse(content)292        except Exception as e:293            return HTMLResponse(f"Error reading file: {e}", status_code=500)294    else:295        return HTMLResponse("File not found", status_code=404)296 297# Download endpoint for agent files298@app.get("/download")299async def download_file(path: str = Query(...)):300    base_dir = AGENT_DIR301    file_path = base_dir / Path(path)302    try:303        file_path.relative_to(base_dir)304    except ValueError:305        return HTMLResponse("Access denied", status_code=403)306    if file_path.exists() and file_path.is_file():307        return FileResponse(308            file_path,309            media_type="application/octet-stream",310            filename=file_path.name311        )312    return HTMLResponse("File not found", status_code=404)313 314# Endpoint to return meta information315@app.get("/meta")316async def get_meta():317    remaining = max(0, int(next_call_time - time.time()))318    key_num = openrouter_key["key_num"]319    return {"key_num": key_num, "remaining": remaining}320 321@app.websocket("/ws")322async def websocket_endpoint(websocket: WebSocket):323    await websocket.accept()324    agent = app.state.agent325    last_position = 0326    try:327        content = agent.log_file.read_text()328        await websocket.send_text(content)329        last_position = len(content)330    except Exception as e:331        logging.error(f"WS init failed: {e}")332    while True:333        await agent.log_update_event.wait()334        agent.log_update_event.clear()335        try:336            with agent.log_file.open("r") as f:337                f.seek(last_position)338                new_content = f.read()339                if new_content:340                    await websocket.send_text(new_content)341                    last_position = f.tell()342        except Exception as e:343            logging.error(f"WS update failed: {e}")344 345if __name__ == "__main__":346    uvicorn.run(app, host="0.0.0.0", port=8000, access_log=False, log_level="warning")347