import os import re import glob import time import functools import sqlite3 import subprocess from typing import List, Tuple, Optional from flask import Flask, request, jsonify, render_template, Response, abort import requests from werkzeug.utils import secure_filename from llama_cpp import Llama from huggingface_hub import hf_hub_download from dotenv import load_dotenv from joserfc import jwt, jws from joserfc.jwk import RSAKey from joserfc.jwt import JWTClaimsRegistry load_dotenv() # Security: rate limiting RATE_LIMIT = 10 # 10 requests per minute REQUEST_COUNTS = {} def rate_limit(func): @functools.wraps(func) def wrapper(*args, **kwargs): client_ip = request.remote_addr now = time.time() if client_ip not in REQUEST_COUNTS: REQUEST_COUNTS[client_ip] = {"count": 0, "timestamp": now} if now - REQUEST_COUNTS[client_ip]["timestamp"] > 60: REQUEST_COUNTS[client_ip] = {"count": 0, "timestamp": now} if REQUEST_COUNTS[client_ip]["count"] >= RATE_LIMIT: abort(429) REQUEST_COUNTS[client_ip]["count"] += 1 return func(*args, **kwargs) return wrapper # --- Google OAuth 2.0 / OIDC Config --- GOOGLE_CLIENT_ID = os.getenv("GOOGLE_CLIENT_ID") def require_oauth(): """Google OIDC token verification decorator implementing OAuth 2.1 specs.""" def oauth_decorator(f): @functools.wraps(f) def oauth_wrapper(*args, **kwargs): auth_header = request.headers.get("Authorization", "") if not auth_header.startswith("Bearer "): return jsonify({"error": "unsupported_token_type", "message": "Missing bearer token."}), 401 token_string = auth_header.split(" ")[1] try: jwks_url = "https://www.googleapis.com/oauth2/v3/certs" jwks_data = requests.get(jwks_url, timeout=5).json() # Correctly access protected header for joserfc obj = jws.extract_compact(token_string.encode()) kid = obj.protected.get("kid") raw_key = next((k for k in jwks_data.get("keys", []) if k.get("kid") == kid), None) if not raw_key: return jsonify({"error": "invalid_key", "message": "Google public key not found."}), 401 public_key = RSAKey.import_key(raw_key) token = jwt.decode(token_string, public_key) claims_registry = JWTClaimsRegistry( iss={"values": ["https://accounts.google.com", "accounts.google.com"]}, aud={"value": GOOGLE_CLIENT_ID} ) claims_registry.validate(token.claims) except Exception as e: print(f"TOKEN VALIDATION ERROR: {str(e)}") return jsonify({"error": "invalid_token", "message": str(e)}), 401 return f(*args, **kwargs) return oauth_wrapper return oauth_decorator # --- Cloud Config --- REPO_ID = "Qwen/Qwen2.5-Coder-0.5B-Instruct-GGUF" FILENAME = "qwen2.5-coder-0.5b-instruct-q4_k_m.gguf" MODEL = REPO_ID os.makedirs("/app/data", exist_ok=True) os.makedirs("/app/models", exist_ok=True) model_path = os.path.join("/app/models", FILENAME) if not os.path.exists(model_path): print(f"Downloading model {FILENAME}...") hf_hub_download(repo_id=REPO_ID, filename=FILENAME, local_dir="/app/models") llm = Llama(model_path=model_path, n_ctx=8192, n_threads=4, n_batch=512, flash_attn=True) DB_PATH = "/app/data/memory.db" PROJECTS_DIR = os.getenv("PROJECTS_DIR", "./projects") TEMPERATURE = float(os.getenv("TEMPERATURE", "0.2")) TOP_P = float(os.getenv("TOP_P", "0.9")) NUM_CTX = int(os.getenv("NUM_CTX", "2048")) SEED = int(os.getenv("SEED", "7")) MAX_FILES_IN_CONTEXT = int(os.getenv("MAX_FILES_IN_CONTEXT", "10")) MAX_FILE_BYTES = int(os.getenv("MAX_FILE_BYTES", str(16 * 1024))) MAX_PROMPT_CHARS = int(os.getenv("MAX_PROMPT_CHARS", str(46000))) ALLOWED_EXTENSIONS = { "py", "ipynb", "js", "ts", "tsx", "jsx", "md", "txt", "json", "yml", "yaml", "html", "css", "toml", "ini", "cfg", "sh", "ps1" } RUNNER_ENABLED = os.getenv("RUNNER_ENABLED", "0") == "1" LINTER_ENABLED = os.getenv("LINTER_ENABLED", "0") == "1" SYSTEM_PROMPT = """ You are a senior AI engineer. 1) Always start your response with '#mode: write|review|explain|discuss|math'. 2) If code is requested, use the 'FILE: path/to/file.ext' format inside markdown blocks. 3) If a general question is asked, provide a clear and direct answer. 4) Be concise and professional. """ app = Flask(__name__, static_folder="../static", template_folder="../templates") conn = sqlite3.connect(DB_PATH, check_same_thread=False) cur = conn.cursor() cur.execute(""" CREATE TABLE IF NOT EXISTS messages ( id INTEGER PRIMARY KEY AUTOINCREMENT, project TEXT, role TEXT, content TEXT, ts REAL ) """) conn.commit() def save_message(project: str, role: str, content: str): cur.execute( "INSERT INTO messages (project, role, content, ts) VALUES (?, ?, ?, ?)", (project, role, content, time.time()) ) conn.commit() def load_recent(project: str, limit: int = 12): cur.execute( "SELECT role, content FROM messages WHERE project=? ORDER BY id DESC LIMIT ?", (project, limit) ) rows = cur.fetchall()[::-1] return [{"role": r[0], "content": r[1]} for r in rows] def project_base_dir(project: str) -> str: base = os.path.abspath(PROJECTS_DIR) path = os.path.abspath(os.path.join(base, project)) if not path.startswith(base + os.sep) and path != base: abort(400, description="Invalid project path") os.makedirs(path, exist_ok=True) return path def _is_allowed_file(path: str) -> bool: return "." in path and path.rsplit(".", 1)[1].lower() in ALLOWED_EXTENSIONS def _sanitize_rel_path(rel_path: str) -> Optional[str]: rel_path = rel_path.strip().replace("\\", "/") rel_path = re.sub(r"^/+", "", rel_path) if not _is_allowed_file(rel_path): return None if ".." in rel_path.split("/"): return None return rel_path def _rank_project_files(file_paths: List[str]) -> List[str]: def key(p: str): try: mtime = os.path.getmtime(p) except OSError: mtime = 0 depth = p.count(os.sep) name_bonus = 0 basename = os.path.basename(p).lower() if basename in {"readme.md", "requirements.txt", "pyproject.toml", "setup.py"}: name_bonus = 10 return (-mtime, depth, -name_bonus) return sorted(file_paths, key=key) def load_project_files_context(project: str) -> str: base_dir = project_base_dir(project) files_data = [] byte_budget = MAX_PROMPT_CHARS preferred_globs = [ "**/*.py", "**/*.ipynb", "**/*.md", "**/*.js", "**/*.ts", "**/*.tsx", "**/*.jsx", "**/*.json", "**/*.yml", "**/*.yaml", "**/*.toml", "**/*.ini", "**/*.html", "**/*.css", "README.md", "requirements.txt", "pyproject.toml", "setup.py" ] candidate_paths = set() for pattern in preferred_globs: candidate_paths.update(glob.glob(os.path.join(base_dir, pattern), recursive=True)) ranked = _rank_project_files([p for p in candidate_paths if os.path.isfile(p)]) count = 0 for path in ranked: if count >= MAX_FILES_IN_CONTEXT: break try: size = os.path.getsize(path) except OSError: continue if size > MAX_FILE_BYTES: continue try: with open(path, "r", encoding="utf-8", errors="ignore") as f: content = f.read() except Exception: continue rel_path = os.path.relpath(path, base_dir) block = f"FILE: {rel_path}\n```\n{content}\n```" if len(block) > byte_budget: break files_data.append(block) byte_budget -= len(block) count += 1 if not files_data: return "No existing project files found." return "\n\n".join(files_data) FILE_BLOCK_RE = re.compile( r"FILE:\s*(?P[^\n\r]+)\s*```(?P[\w.+-]+)?\s*\n(?P.*?)```", re.DOTALL ) CODE_FENCE_RE = re.compile( r"```(?P[\w.+-]*)\s*\n(?P.*?)```", re.DOTALL ) def _ensure_mode_header(text: str, default_mode: str = "write") -> str: first = text.strip().splitlines()[0].strip() if text.strip().splitlines() else "" if not re.search(r"^#mode:\s*(write|review|explain)\s*$", first, re.IGNORECASE): text = f"#mode: {default_mode}\n\n" + text return text def _ensure_plan_section(text: str) -> str: if re.search(r"(?im)^\s*plan\s*$", text): return text preface = ( "Plan\n" "- Outline steps briefly.\n" "- Write complete code using the Multi-file format.\n" "- Add a tiny MWE when appropriate.\n" "- Include Self-Check at the end.\n\n" ) return text if "Plan" in text[:400] else preface + text def _wrap_lonely_fence_as_file(text: str) -> str: if FILE_BLOCK_RE.search(text): return text m = CODE_FENCE_RE.search(text) if not m: return text lang = (m.group("lang") or "text").lower() default_map = { "python": "scratch/main.py", "py": "scratch/main.py", "javascript": "scratch/index.js", "js": "scratch/index.js", "typescript": "scratch/index.ts", "ts": "scratch/index.ts", "json": "scratch/data.json", "html": "scratch/index.html", "css": "scratch/styles.css", "md": "scratch/README.md", } rel = default_map.get(lang, "scratch/snippet.txt") code = m.group("code") file_block = f"FILE: {rel}\n```{lang}\n{code}\n```" start, end = m.span() return text[:start] + file_block + text[end:] def enforce_response_contract(text: str, default_mode: str = "write") -> str: text = _ensure_mode_header(text, default_mode=default_mode) text = _ensure_plan_section(text) text = _wrap_lonely_fence_as_file(text) text = re.sub(r"```(\s*\n)", "```text\1", text) return text def save_generated_files(project: str, assistant_text: str) -> List[str]: base_dir = project_base_dir(project) matches = list(FILE_BLOCK_RE.finditer(assistant_text)) saved: List[str] = [] for m in matches: rel_path_raw = m.group("path") lang = (m.group("lang") or "").strip().lower() code = m.group("code") rel_path = _sanitize_rel_path(rel_path_raw) if not rel_path: continue abs_path = os.path.abspath(os.path.join(base_dir, rel_path)) if not abs_path.startswith(base_dir + os.sep) and abs_path != base_dir: continue os.makedirs(os.path.dirname(abs_path), exist_ok=True) if not lang: ext = rel_path.rsplit(".", 1)[-1].lower() lang = ext with open(abs_path, "w", encoding="utf-8") as f: f.write(code.strip()) saved.append(rel_path) return saved def parse_user_mode(text: str) -> str: m = re.search(r"#mode:\s*(write|review|explain|discuss|math)", text, re.IGNORECASE) if m: return m.group(1).lower() if re.search(r"write|code|create|build|script|function|generate|program|app", text, re.IGNORECASE): return "write" if re.search(r"review|critique|improve|refactor|fix|bug|error", text, re.IGNORECASE): return "review" if re.search(r"explain|walk me through|how does|what does.*mean", text, re.IGNORECASE): return "explain" return "discuss" def build_full_prompt(project: str, user_text: str, search_results: List[str] = None) -> Tuple[str, str]: hist = load_recent(project, limit=8) hist_lines = [f"{m['role'].upper()}: {m['content']}" for m in hist] files_context = load_project_files_context(project) mode = parse_user_mode(user_text) search_context = "" if search_results: search_context = "\n--- Web Search Results ---\n" + "\n".join(search_results) if mode in ["write", "review"]: planning_instructions = f""" IMPORTANT: - Start with '#mode: {mode}'. - ONLY if you are generating code: Write a 'Plan' section and use the Multi-file format. - If this is a general question: Answer directly and ignore the 'Plan' requirement. - End with a 'Self-Check' if code was written. """ else: planning_instructions = f""" IMPORTANT: - Start with '#mode: {mode}'. - Use the provided search results to answer the user's question directly. - Be concise and do not use coding formats. """ sections = [ SYSTEM_PROMPT.strip(), "\n--- Session Settings ---\n", f"Model: {MODEL}\nTemperature: {TEMPERATURE}\n", search_context, "\n--- Project Context ---\n", files_context, "\n--- Conversation ---\n", "\n".join(hist_lines), "\n--- New Request ---\n", f"USER: {user_text}\n{planning_instructions}\nASSISTANT:" ] prompt = "\n".join(s for s in sections if s and s.strip()) if len(prompt) > MAX_PROMPT_CHARS: prompt = prompt[-MAX_PROMPT_CHARS:] return prompt, mode # -------------- Routes -------------- @app.route("/") @rate_limit def index(): return render_template( "index.html", google_client_id=GOOGLE_CLIENT_ID ) @app.route("/protected") @require_oauth() def protected_resource(): return "This is a protected resource!" @app.route("/history/", methods=["GET"]) def get_history(project): try: conn = sqlite3.connect(DB_PATH) cur = conn.cursor() cur.execute("SELECT role, content FROM messages WHERE project=? ORDER BY id ASC", (project,)) rows = cur.fetchall() conn.close() history_list = [{"role": r[0], "content": r[1]} for r in rows] return jsonify({"history": history_list}) except Exception as e: print(f"Database Error: {e}") return jsonify({"error": str(e)}), 500 @app.route("/projects") def get_projects(): cur.execute("SELECT DISTINCT project FROM messages ORDER BY project ASC") rows = cur.fetchall() return jsonify([r[0] for r in rows]) @app.route("/settings", methods=["GET", "POST"]) @require_oauth() def settings(): global MODEL, TEMPERATURE, TOP_P, NUM_CTX, SEED if request.method == "POST": data = request.json or {} MODEL = data.get("model", MODEL) TEMPERATURE = float(data.get("temperature", TEMPERATURE)) TOP_P = float(data.get("top_p", TOP_P)) NUM_CTX = int(data.get("num_ctx", NUM_CTX)) SEED = int(data.get("seed", SEED)) return jsonify({ "model": MODEL, "temperature": TEMPERATURE, "top_p": TOP_P, "num_ctx": NUM_CTX, "seed": SEED }) @app.route("/add_project", methods=["POST"]) def add_project(): data = request.json or {} project = (data.get("project") or "").strip() if not project: return jsonify({"error": "empty project name"}), 400 project_base_dir(project) save_message(project, "system", f"Project {project} created.") return jsonify({"status": "ok", "project": project}) @app.route("/chat", methods=["POST"]) @require_oauth() def chat(): data = request.json or {} project = data.get("project", "default") user_text = (data.get("message") or "").strip() if not user_text: return jsonify({"error": "empty message"}), 400 save_message(project, "user", user_text) full_prompt, mode = build_full_prompt(project, user_text) try: resp = llm( full_prompt, max_tokens=512, temperature=TEMPERATURE, top_p=TOP_P, stop=["\nUSER:", "\nSYSTEM:"], echo=False ) assistant_text_raw = resp["choices"][0]["text"] assistant_text = enforce_response_contract(assistant_text_raw, default_mode=mode) except Exception as e: return jsonify({"error": str(e)}), 500 save_message(project, "assistant", assistant_text) saved_files = save_generated_files(project, assistant_text) return jsonify({"response": assistant_text, "saved_files": saved_files}) @app.route("/stream", methods=["POST"]) @require_oauth() @rate_limit def stream(): data = request.json or {} project = data.get("project", "default") user_text = (data.get("message") or "").strip() save_message(project, "user", user_text) full_prompt, mode = build_full_prompt(project, user_text) def generate(): yield "data: \n\n" try: stream_res = llm( full_prompt, max_types=1024, temperature=0.7, stream=True, stop=["USER:", "ASSISTANT:"] ) for chunk in stream_res: token = chunk.get("choices", [{}])[0].get("text", "") if token: safe_token = token.replace("\n", "\\n").replace("\r", "") yield f"data: {safe_token}\n\n" yield "data: [DONE]\n\n" except Exception as e: print(f"STREAM ERROR: {e}") yield f"data: ERROR: {str(e)}\n\n" resp = Response(generate(), mimetype="text/event-stream") resp.headers["X-Accel-Buffering"] = "no" resp.headers["Cache-Control"] = "no-cache" return resp @app.route("/search_web", methods=["POST"]) @require_oauth() @rate_limit def search_web(): data = request.json or {} query = (data.get("query") or "").strip() TAVILY_API_KEY = os.getenv("TAVILY_API_KEY") payload = {"api_key": TAVILY_API_KEY, "query": query, "search_depth": "basic", "max_results": 3} try: response = requests.post("https://api.tavily.com/search", json=payload, timeout=10) tavily_data = response.json() results = [f"{r['title']}: {r['content']}" for r in tavily_data.get("results", [])] return jsonify({"results": results}) except Exception as e: return jsonify({"error": str(e)}), 500 def allowed_file(filename): return "." in filename and filename.rsplit(".", 1)[1].lower() in ALLOWED_EXTENSIONS @app.route("/upload_file/", methods=["POST"]) @require_oauth() @rate_limit def upload_file(project): if "file" not in request.files: return jsonify({"error": "no file part"}), 400 file = request.files["file"] if file.filename == "": return jsonify({"error": "no selected file"}), 400 if file and allowed_file(file.filename): filename = secure_filename(file.filename) save_path = os.path.join(project_base_dir(project), filename) os.makedirs(os.path.dirname(save_path), exist_ok=True) file.save(save_path) save_message(project, "system", f"File uploaded: {filename}") return jsonify({"status": "ok", "filename": filename}) else: return jsonify({"error": "file type not allowed"}), 400 @app.route("/delete_project", methods=["POST"]) @require_oauth() def delete_project(): data = request.json or {} project = (data.get("project") or "").strip() if not project: return jsonify({"error": "empty project name"}), 400 cur.execute("DELETE FROM messages WHERE project=?", (project,)) conn.commit() project_dir = os.path.join(PROJECTS_DIR, project) if os.path.exists(project_dir): import shutil shutil.rmtree(project_dir) return jsonify({"status": "ok", "project": project}) @app.route("/cancel", methods=["POST"]) def cancel(): return jsonify({"status": "Session reset requested", "note": "Inference is self-contained."}) def _run_cmd(cmd: List[str], cwd: Optional[str] = None, timeout: int = 20) -> Tuple[int, str, str]: try: p = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True, timeout=timeout, check=False) return p.returncode, p.stdout, p.stderr except Exception as e: return 1, "", str(e) @app.route("/run/", methods=["POST"]) @require_oauth() @rate_limit def run_project(project): if not RUNNER_ENABLED: return jsonify({"error": "runner disabled; set RUNNER_ENABLED=1"}), 400 data = request.json or {} entry = (data.get("entry") or "main.py").strip() base = project_base_dir(project) if not _is_allowed_file(entry): return jsonify({"error": "disallowed entry point"}), 400 path = os.path.join(base, entry) if not os.path.exists(path): return jsonify({"error": f"missing entry: {entry}"}), 404 code, out, err = _run_cmd(["python", entry], cwd=base, timeout=60) return jsonify({"code": code, "stdout": out, "stderr": err}) @app.route("/lint/", methods=["POST"]) @require_oauth() @rate_limit def lint(project): if not LINTER_ENABLED: return jsonify({"error": "linter disabled; set LINTER_ENABLED=1"}), 400 base = project_base_dir(project) try: import shutil as _shutil has_ruff = _shutil.which("ruff") is not None except Exception: has_ruff = False if has_ruff: code, out, err = _run_cmd(["ruff", "."], cwd=base, timeout=60) else: code, out, err = _run_cmd(["python", "-m", "pyflakes", "."], cwd=base, timeout=60) return jsonify({"code": code, "stdout": out, "stderr": err}) if __name__ == "__main__": os.makedirs(PROJECTS_DIR, exist_ok=True) app.run(host="0.0.0.0", port=5000, debug=True)