File size: 5,467 Bytes
4554903
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Model access for Rivet.

Two backends behind one interface:

  OllamaClient       — HTTP to an Ollama server. Knowledge packs are
                       injected as a system-context block (Ollama's API
                       exposes no KV-cache handle, so this is the honest
                       limit of that transport).
  TransformersClient — local HF transformers. Knowledge packs are
                       injected as precomputed KV cache blocks via
                       pharos.kv_injector — true zero-prompt-token
                       injection, the real Pharos pipeline.

Which backend runs is a config decision (`model.backend`), not a code
change. Chips only ever see `ModelClient.generate()`.
"""

import json
import urllib.error
import urllib.request
from dataclasses import dataclass


@dataclass
class ModelReply:
    text: str
    ok: bool
    backend: str
    knowledge_injected: str = "none"  # none | system_prompt | kv_cache
    error: str = ""


class ModelClient:
    """Interface. Use OllamaClient or TransformersClient."""

    def generate(self, prompt: str, system: str = "", knowledge: str = "",
                 max_tokens: int = 1024) -> ModelReply:
        raise NotImplementedError


class OllamaClient(ModelClient):
    def __init__(self, base_url: str = "http://localhost:11434",
                 model: str = "qwen2.5-coder:32b",
                 temperature: float = 0.3, num_ctx: int = 32768,
                 timeout: int = 180):
        self.base_url = base_url.rstrip("/")
        self.model = model
        self.temperature = temperature
        self.num_ctx = num_ctx
        self.timeout = timeout

    def generate(self, prompt: str, system: str = "", knowledge: str = "",
                 max_tokens: int = 1024) -> ModelReply:
        full_system = system
        injected = "none"
        if knowledge:
            full_system = (
                f"{system}\n\n# INJECTED KNOWLEDGE (Pharos)\n"
                f"You have access to the following knowledge:\n{knowledge}"
            ).strip()
            injected = "system_prompt"

        payload = {
            "model": self.model,
            "prompt": prompt,
            "system": full_system,
            "stream": False,
            "options": {
                "temperature": self.temperature,
                "num_ctx": self.num_ctx,
                "num_predict": max_tokens,
            },
        }
        req = urllib.request.Request(
            f"{self.base_url}/api/generate",
            data=json.dumps(payload).encode(),
            headers={"Content-Type": "application/json"},
        )
        try:
            with urllib.request.urlopen(req, timeout=self.timeout) as resp:
                data = json.loads(resp.read())
            return ModelReply(
                text=data.get("response", ""), ok=True,
                backend=f"ollama:{self.model}", knowledge_injected=injected,
            )
        except (urllib.error.URLError, TimeoutError, OSError) as exc:
            return ModelReply(
                text="", ok=False, backend=f"ollama:{self.model}",
                error=f"Ollama unreachable at {self.base_url}: {exc}",
            )


class TransformersClient(ModelClient):
    """Local transformers backend with real KV-cache pack injection.

    Lazy-loads torch/transformers on first generate() so importing this
    module never requires them. Pack KV blocks come from
    pharos.kv_injector.KVPackEncoder (precomputed once per pack).
    """

    def __init__(self, model_path: str, device: str = "auto",
                 pack_cache_dir: str = ""):
        self.model_path = model_path
        self.device = device
        self.pack_cache_dir = pack_cache_dir
        self._injector = None

    def _ensure_loaded(self):
        if self._injector is None:
            from pharos.kv_injector import KVInjector
            self._injector = KVInjector(
                self.model_path, device=self.device,
                cache_dir=self.pack_cache_dir,
            )
        return self._injector

    def generate(self, prompt: str, system: str = "", knowledge: str = "",
                 max_tokens: int = 1024) -> ModelReply:
        try:
            injector = self._ensure_loaded()
        except ImportError as exc:
            return ModelReply(
                text="", ok=False, backend="transformers",
                error=f"torch/transformers not available: {exc}",
            )
        text = injector.generate(
            prompt, system=system, knowledge=knowledge, max_tokens=max_tokens,
        )
        return ModelReply(
            text=text, ok=True,
            backend=f"transformers:{self.model_path}",
            knowledge_injected="kv_cache" if knowledge else "none",
        )


def build_client(cfg: dict) -> ModelClient:
    """Construct the configured backend from the `model:` config section."""
    backend = cfg.get("backend", "ollama")
    if backend == "transformers":
        return TransformersClient(
            model_path=cfg.get("model_path", cfg.get("name", "")),
            device=cfg.get("device", "auto"),
            pack_cache_dir=cfg.get("pack_cache_dir", ""),
        )
    return OllamaClient(
        base_url=cfg.get("base_url", "http://localhost:11434"),
        model=cfg.get("name", "qwen2.5-coder:32b"),
        temperature=cfg.get("temperature", 0.3),
        num_ctx=cfg.get("num_ctx", 32768),
        timeout=cfg.get("timeout_seconds", 180),
    )