| 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, |
| ) |
|
|