from __future__ import annotations import ast import ipaddress import operator import os import re import socket import threading from datetime import datetime from pathlib import Path from typing import Any from urllib.parse import urljoin, urlparse from dotenv import load_dotenv load_dotenv() PERSISTENT_ROOT = Path(os.getenv("PERSISTENT_ROOT", "/data")).expanduser() try: PERSISTENT_ROOT.mkdir(parents=True, exist_ok=True) except OSError: PERSISTENT_ROOT = Path("data") PERSISTENT_ROOT.mkdir(parents=True, exist_ok=True) os.environ.setdefault("HF_HOME", str(PERSISTENT_ROOT / "huggingface")) os.environ.setdefault("HF_HUB_CACHE", str(PERSISTENT_ROOT / "huggingface" / "hub")) import gradio as gr import requests from bs4 import BeautifulSoup from ddgs import DDGS from huggingface_hub import hf_hub_download from llama_cpp import Llama from smolagents import ChatMessage, CodeAgent, Model, Tool MODEL_REPO = os.getenv("MODEL_REPO", "NANI-Nithin/K2-Horizon-0.9B-GGUF") MODEL_FILE = os.getenv("MODEL_FILE", "K2-Horizon-0.9B-Q4_K_M.gguf") MODEL_DIR = Path(os.getenv("MODEL_DIR", str(PERSISTENT_ROOT / "models"))).expanduser() def env_int(name: str, default: int) -> int: try: return int(os.getenv(name, str(default))) except ValueError: return default def env_float(name: str, default: float) -> float: try: return float(os.getenv(name, str(default))) except ValueError: return default def content_to_text(content: Any) -> str: if isinstance(content, str): return content if isinstance(content, list): parts: list[str] = [] for item in content: if isinstance(item, dict): parts.append(str(item.get("text", item.get("content", "")))) else: parts.append(str(item)) return "\n".join(part for part in parts if part) return str(content or "") class LlamaCppModel(Model): def __init__(self, llama: Llama, max_tokens: int, temperature: float) -> None: super().__init__() self.llama = llama self.max_tokens = max_tokens self.temperature = temperature @staticmethod def _normalize_messages(messages: list[Any]) -> list[dict[str, str]]: normalized: list[dict[str, str]] = [] for message in messages: if isinstance(message, dict): role = message.get("role", "user") content = message.get("content", "") else: role = getattr(message, "role", "user") content = getattr(message, "content", "") if hasattr(role, "value"): role = role.value role = str(role).lower() if role not in {"system", "user", "assistant"}: role = "user" normalized.append({"role": role, "content": content_to_text(content)}) return normalized def generate( self, messages: list[Any], stop_sequences: list[str] | None = None, response_format: dict[str, Any] | None = None, tools_to_call_from: list[Tool] | None = None, **kwargs: Any, ) -> ChatMessage: del response_format, tools_to_call_from result = self.llama.create_chat_completion( messages=self._normalize_messages(messages), max_tokens=int(kwargs.get("max_tokens", self.max_tokens)), temperature=float(kwargs.get("temperature", self.temperature)), top_p=float(kwargs.get("top_p", 0.9)), repeat_penalty=float(kwargs.get("repeat_penalty", 1.1)), stop=stop_sequences or None, ) content = result["choices"][0]["message"].get("content", "") return ChatMessage(role="assistant", content=content) def __call__(self, messages: list[Any], **kwargs: Any) -> ChatMessage: return self.generate(messages, **kwargs) def direct_chat(self, messages: list[dict[str, str]]) -> str: result = self.llama.create_chat_completion( messages=messages, max_tokens=self.max_tokens, temperature=self.temperature, top_p=0.9, repeat_penalty=1.1, ) return str(result["choices"][0]["message"].get("content", "")).strip() class DuckDuckGoSearchTool(Tool): name = "web_search" description = "Search the public web with DuckDuckGo. Use it for current facts and external information." inputs = { "query": {"type": "string", "description": "A focused web search query."}, "max_results": { "type": "integer", "description": "Number of results from 1 to 8.", "nullable": True, }, } output_type = "string" def forward(self, query: str, max_results: int | None = None) -> str: limit = max(1, min(int(max_results or 5), 8)) results = list(DDGS().text(query, max_results=limit)) if not results: return "No search results found." rows = [] for index, item in enumerate(results, 1): title = item.get("title", "Untitled") url = item.get("href", item.get("url", "")) body = item.get("body", "") rows.append(f"{index}. {title}\nURL: {url}\nSnippet: {body}") return "\n\n".join(rows) def ensure_public_url(url: str) -> str: parsed = urlparse(url) if parsed.scheme not in {"http", "https"} or not parsed.hostname: raise ValueError("Only public http/https URLs are allowed.") addresses = socket.getaddrinfo(parsed.hostname, parsed.port or 80, proto=socket.IPPROTO_TCP) for address in addresses: ip = ipaddress.ip_address(address[4][0]) if not ip.is_global: raise ValueError("Private, loopback, and local network addresses are blocked.") return url class ReadWebpageTool(Tool): name = "read_webpage" description = "Download and extract readable text from a public web page URL." inputs = { "url": {"type": "string", "description": "The full public http or https URL."}, } output_type = "string" def forward(self, url: str) -> str: current_url = url response = None for _ in range(4): safe_url = ensure_public_url(current_url) response = requests.get( safe_url, timeout=12, allow_redirects=False, headers={"User-Agent": "Mozilla/5.0 (compatible; K2-Horizon-Agent/1.0)"}, ) if response.status_code not in {301, 302, 303, 307, 308}: break location = response.headers.get("location") if not location: break current_url = urljoin(current_url, location) assert response is not None response.raise_for_status() content_type = response.headers.get("content-type", "") if "text/html" not in content_type and "text/plain" not in content_type: return f"Unsupported content type: {content_type}" soup = BeautifulSoup(response.text[:2_000_000], "html.parser") for node in soup(["script", "style", "noscript", "svg"]): node.decompose() text = re.sub(r"\n{3,}", "\n\n", soup.get_text("\n", strip=True)) return text[:12_000] or "No readable text found." _BINARY_OPERATORS = { ast.Add: operator.add, ast.Sub: operator.sub, ast.Mult: operator.mul, ast.Div: operator.truediv, ast.FloorDiv: operator.floordiv, ast.Mod: operator.mod, ast.Pow: operator.pow, } _UNARY_OPERATORS = {ast.UAdd: operator.pos, ast.USub: operator.neg} def evaluate_expression(node: ast.AST) -> float | int: if isinstance(node, ast.Expression): return evaluate_expression(node.body) if isinstance(node, ast.Constant) and isinstance(node.value, (int, float)): return node.value if isinstance(node, ast.BinOp) and type(node.op) in _BINARY_OPERATORS: left = evaluate_expression(node.left) right = evaluate_expression(node.right) if isinstance(node.op, ast.Pow) and abs(right) > 100: raise ValueError("Exponent is too large.") return _BINARY_OPERATORS[type(node.op)](left, right) if isinstance(node, ast.UnaryOp) and type(node.op) in _UNARY_OPERATORS: return _UNARY_OPERATORS[type(node.op)](evaluate_expression(node.operand)) raise ValueError("Only numeric arithmetic is supported.") class CalculatorTool(Tool): name = "calculator" description = "Safely evaluate a numeric arithmetic expression." inputs = {"expression": {"type": "string", "description": "Arithmetic expression to evaluate."}} output_type = "string" def forward(self, expression: str) -> str: if len(expression) > 200: raise ValueError("Expression is too long.") value = evaluate_expression(ast.parse(expression, mode="eval")) return str(value) class CurrentTimeTool(Tool): name = "current_time" description = "Get the current system date and time, including timezone." inputs = {} output_type = "string" def forward(self) -> str: return datetime.now().astimezone().isoformat(timespec="seconds") class AppRuntime: def __init__(self) -> None: self.lock = threading.Lock() self.model: LlamaCppModel | None = None self.agent: CodeAgent | None = None self.model_path: Path | None = None def load(self) -> None: if self.model is not None: return with self.lock: if self.model is not None: return MODEL_DIR.mkdir(parents=True, exist_ok=True) local_path = hf_hub_download( repo_id=MODEL_REPO, filename=MODEL_FILE, local_dir=str(MODEL_DIR), ) self.model_path = Path(local_path) llama = Llama( model_path=str(self.model_path), n_ctx=env_int("N_CTX", 4096), n_threads=env_int("N_THREADS", max(1, (os.cpu_count() or 4) - 1)), n_threads_batch=env_int("N_THREADS_BATCH", os.cpu_count() or 4), n_batch=env_int("N_BATCH", 256), n_gpu_layers=0, use_mmap=True, verbose=os.getenv("LLAMA_VERBOSE", "0") == "1", ) self.model = LlamaCppModel( llama=llama, max_tokens=env_int("MAX_NEW_TOKENS", 700), temperature=env_float("TEMPERATURE", 0.2), ) self.agent = CodeAgent( tools=[ DuckDuckGoSearchTool(), ReadWebpageTool(), CalculatorTool(), CurrentTimeTool(), ], model=self.model, max_steps=env_int("AGENT_MAX_STEPS", 5), add_base_tools=False, additional_authorized_imports=[], code_block_tags="markdown", ) def reply(self, message: str, history: list[dict[str, str]], use_tools: bool) -> str: self.load() assert self.model is not None if use_tools: assert self.agent is not None transcript = "\n".join( f"{item.get('role', 'user')}: {content_to_text(item.get('content', ''))}" for item in history[-6:] if item.get("role") in {"user", "assistant"} ) task = message if transcript: task = f"Conversation context:\n{transcript}\n\nCurrent user request:\n{message}" return str(self.agent.run(task, reset=True)).strip() system = { "role": "system", "content": "You are K2 Horizon, a concise and helpful local assistant.", } context = [system] for item in history[-10:]: if item.get("role") in {"user", "assistant"}: context.append({"role": item["role"], "content": content_to_text(item.get("content", ""))}) context.append({"role": "user", "content": message}) return self.model.direct_chat(context) runtime = AppRuntime() def respond(message: str, history: list[dict[str, str]], use_tools: bool): if not message.strip(): return "", history updated = history + [{"role": "user", "content": message}] try: answer = runtime.reply(message.strip(), history, use_tools) except Exception as exc: answer = f"Error: {type(exc).__name__}: {exc}" updated.append({"role": "assistant", "content": answer}) return "", updated CSS = """ .gradio-container { max-width: 860px !important; margin: 0 auto !important; } #app-shell { min-height: 100vh; padding: 32px 12px 20px; } #title { text-align: center; margin-bottom: 2px; } #subtitle { text-align: center; color: var(--body-text-color-subdued); margin-bottom: 18px; } #chat { border: 1px solid var(--border-color-primary); border-radius: 18px; overflow: hidden; } #composer { gap: 10px; align-items: stretch; margin-top: 12px; } #prompt textarea { border-radius: 14px !important; } #send { min-width: 92px; border-radius: 14px !important; } #controls { align-items: center; margin-top: 8px; } #note { color: var(--body-text-color-subdued); font-size: 12px; text-align: right; } footer { display: none !important; } """ THEME = gr.themes.Base(primary_hue="slate", neutral_hue="slate") with gr.Blocks( title="K2 Horizon", ) as demo: with gr.Column(elem_id="app-shell"): gr.Markdown("# K2 Horizon", elem_id="title") gr.Markdown("Private CPU inference with optional web tools", elem_id="subtitle") chatbot = gr.Chatbot( height=570, buttons=["copy"], allow_tags=False, placeholder="Ask anything", elem_id="chat", ) with gr.Row(elem_id="composer"): prompt = gr.Textbox( placeholder="Message K2 Horizon…", show_label=False, scale=9, lines=1, max_lines=5, elem_id="prompt", ) send = gr.Button("Send", variant="primary", scale=1, elem_id="send") with gr.Row(elem_id="controls"): use_tools = gr.Checkbox(value=True, label="Web tools", scale=1) clear = gr.Button("Clear", variant="secondary", size="sm", scale=0) gr.Markdown("Model loads on the first message", elem_id="note") send.click(respond, [prompt, chatbot, use_tools], [prompt, chatbot]) prompt.submit(respond, [prompt, chatbot, use_tools], [prompt, chatbot]) clear.click(lambda: ("", []), outputs=[prompt, chatbot], queue=False) if __name__ == "__main__": demo.queue(default_concurrency_limit=1).launch( server_name=os.getenv("GRADIO_SERVER_NAME", "0.0.0.0"), server_port=env_int("GRADIO_SERVER_PORT", 7860), share=os.getenv("GRADIO_SHARE", "0") == "1", show_error=True, ssr_mode=False, footer_links=[], theme=THEME, css=CSS, )