Felipe97/llama-cpp-compiled
01.2k
1#!/usr/bin/env python32# -*- coding: utf-8 -*-3 4# type: ignore[reportUnusedImport]5 6import subprocess7import os8 9TMP_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "tmp")10import re11import json12from json import JSONDecodeError13import sys14import requests15import time16from concurrent.futures import ThreadPoolExecutor, as_completed17from typing import (18 Any,19 Callable,20 ContextManager,21 Iterable,22 Iterator,23 List,24 Literal,25 Tuple,26 Set,27)28from re import RegexFlag29import wget30 31 32DEFAULT_HTTP_TIMEOUT = 6033 34# per-request timeout, a hung server fails the test instead of stalling the CI for hours35DEFAULT_REQUEST_TIMEOUT = 60036 37 38class ServerResponse:39 headers: dict40 status_code: int41 body: dict | Any42 43 44class ServerError(Exception):45 def __init__(self, code, body):46 self.code = code47 self.body = body48 49 50class ServerProcess:51 # default options52 debug: bool = False53 server_port: int = 808054 server_host: str = "127.0.0.1"55 model_hf_repo: str | None = "ggml-org/models"56 model_hf_file: str | None = "tinyllamas/stories260K.gguf"57 model_alias: str = "tinyllama-2"58 temperature: float = 0.859 seed: int = 4260 offline: bool = False61 62 # custom options63 model_alias: str | None = None64 model_tags: str | None = None65 model_url: str | None = None66 model_file: str | None = None67 model_draft: str | None = None68 n_threads: int | None = None69 n_gpu_layer: int | None = None70 n_batch: int | None = None71 n_ubatch: int | None = None72 n_ctx: int | None = None73 n_ga: int | None = None74 n_ga_w: int | None = None75 n_predict: int | None = None76 n_prompts: int | None = 077 slot_save_path: str | None = None78 id_slot: int | None = None79 cache_prompt: bool | None = None80 n_slots: int | None = None81 ctk: str | None = None82 ctv: str | None = None83 fa: str | None = None84 server_continuous_batching: bool | None = False85 server_embeddings: bool | None = False86 server_reranking: bool | None = False87 server_metrics: bool | None = False88 kv_unified: bool | None = False89 swa_full: bool | None = False90 server_slots: bool | None = False91 pooling: str | None = None92 api_key: str | None = None93 models_dir: str | None = None94 models_max: int | None = None95 models_preset: str | None = None96 no_models_autoload: bool | None = None97 lora_files: List[str] | None = None98 enable_ctx_shift: int | None = False99 spec_type: str | None = None100 spec_draft_n_min: int | None = None101 spec_draft_n_max: int | None = None102 spec_synth_len: float | None = None103 spec_synth_rates: List[float] | None = None104 no_ui: bool | None = None105 jinja: bool | None = None106 reasoning_format: Literal['deepseek', 'none', 'nothink'] | None = None107 reasoning: Literal['on', 'off', 'auto'] | None = None108 chat_template: str | None = None109 chat_template_file: str | None = None110 server_path: str | None = None111 mmproj_url: str | None = None112 no_mmproj: bool | None = None113 media_path: str | None = None114 sleep_idle_seconds: int | None = None115 cache_ram: int | None = None116 no_cache_idle_slots: bool = False117 log_path: str | None = None118 ui_mcp_proxy: bool = False119 backend_sampling: bool = False120 gcp_compat: bool = False121 server_tools: str | None = None122 server_tools_runtime: str | None = None123 mcp_servers_config: str | None = None124 mcp_servers_json: str | None = None125 cors_origins: str | None = None126 127 # session variables128 process: subprocess.Popen | None = None129 130 def __init__(self):131 if "N_GPU_LAYERS" in os.environ:132 self.n_gpu_layer = int(os.environ["N_GPU_LAYERS"])133 if "DEBUG" in os.environ:134 self.debug = True135 if "PORT" in os.environ:136 self.server_port = int(os.environ["PORT"])137 self.external_server = "DEBUG_EXTERNAL" in os.environ138 139 def start(self, timeout_seconds: int = DEFAULT_HTTP_TIMEOUT) -> None:140 env = {141 **os.environ,142 "LLAMA_SERVER_DEBUG_FAKE_TIMING": "1",143 }144 if "LLAMA_CACHE" not in os.environ:145 env["LLAMA_CACHE"] = "tmp"146 if self.external_server:147 print(f"[external_server]: Assuming external server running on {self.server_host}:{self.server_port}")148 return149 if self.server_path is not None:150 server_path = self.server_path151 elif "LLAMA_SERVER_BIN_PATH" in os.environ:152 server_path = os.environ["LLAMA_SERVER_BIN_PATH"]153 elif os.name == "nt":154 server_path = "../../../build/bin/Release/llama-server.exe"155 else:156 server_path = "../../../build/bin/llama-server"157 server_args = [158 "--host",159 self.server_host,160 "--port",161 self.server_port,162 "--temp",163 self.temperature,164 "--seed",165 self.seed,166 ]167 if self.offline:168 server_args.append("--offline")169 if self.model_file:170 server_args.extend(["--model", self.model_file])171 if self.model_url:172 server_args.extend(["--model-url", self.model_url])173 if self.model_draft:174 server_args.extend(["--model-draft", self.model_draft])175 if self.model_hf_repo:176 server_args.extend(["--hf-repo", self.model_hf_repo])177 if self.model_hf_file:178 server_args.extend(["--hf-file", self.model_hf_file])179 if self.models_dir:180 server_args.extend(["--models-dir", self.models_dir])181 if self.models_max is not None:182 server_args.extend(["--models-max", self.models_max])183 if self.models_preset:184 server_args.extend(["--models-preset", self.models_preset])185 if self.cors_origins:186 server_args.extend(["--cors-origins", self.cors_origins])187 if self.n_batch:188 server_args.extend(["--batch-size", self.n_batch])189 if self.n_ubatch:190 server_args.extend(["--ubatch-size", self.n_ubatch])191 if self.n_threads:192 server_args.extend(["--threads", self.n_threads])193 if self.n_gpu_layer:194 server_args.extend(["--n-gpu-layers", self.n_gpu_layer])195 if self.server_continuous_batching:196 server_args.append("--cont-batching")197 if self.server_embeddings:198 server_args.append("--embedding")199 if self.server_reranking:200 server_args.append("--reranking")201 if self.server_metrics:202 server_args.append("--metrics")203 if self.kv_unified:204 server_args.append("--kv-unified")205 if self.swa_full:206 server_args.append("--swa-full")207 if self.server_slots:208 server_args.append("--slots")209 else:210 server_args.append("--no-slots")211 if self.pooling:212 server_args.extend(["--pooling", self.pooling])213 if self.model_alias:214 server_args.extend(["--alias", self.model_alias])215 if self.model_tags:216 server_args.extend(["--tags", self.model_tags])217 if self.n_ctx:218 server_args.extend(["--ctx-size", self.n_ctx])219 if self.n_slots:220 server_args.extend(["--parallel", self.n_slots])221 if self.ctk:222 server_args.extend(["-ctk", self.ctk])223 if self.ctv:224 server_args.extend(["-ctv", self.ctv])225 if self.fa is not None:226 server_args.extend(["-fa", self.fa])227 if self.n_predict:228 server_args.extend(["--n-predict", self.n_predict])229 if self.slot_save_path:230 server_args.extend(["--slot-save-path", self.slot_save_path])231 if self.n_ga:232 server_args.extend(["--grp-attn-n", self.n_ga])233 if self.n_ga_w:234 server_args.extend(["--grp-attn-w", self.n_ga_w])235 if self.debug:236 server_args.append("--verbose")237 if self.lora_files:238 for lora_file in self.lora_files:239 server_args.extend(["--lora", lora_file])240 if self.enable_ctx_shift:241 server_args.append("--context-shift")242 if self.spec_type:243 server_args.extend(["--spec-type", self.spec_type])244 if self.api_key:245 server_args.extend(["--api-key", self.api_key])246 if self.spec_draft_n_max:247 server_args.extend(["--spec-draft-n-max", self.spec_draft_n_max])248 if self.spec_draft_n_min:249 server_args.extend(["--spec-draft-n-min", self.spec_draft_n_min])250 if self.spec_synth_len is not None:251 server_args.extend(["--spec-synth-len", self.spec_synth_len])252 if self.spec_synth_rates is not None:253 rates = ",".join(str(rate) for rate in self.spec_synth_rates)254 server_args.extend(["--spec-synth-rates", rates])255 if self.no_ui:256 server_args.append("--no-ui")257 if self.no_models_autoload:258 server_args.append("--no-models-autoload")259 if self.jinja:260 server_args.append("--jinja")261 else:262 server_args.append("--no-jinja")263 if self.reasoning_format is not None:264 server_args.extend(("--reasoning-format", self.reasoning_format))265 if self.reasoning is not None:266 server_args.extend(("--reasoning", self.reasoning))267 if self.chat_template:268 server_args.extend(["--chat-template", self.chat_template])269 if self.chat_template_file:270 server_args.extend(["--chat-template-file", self.chat_template_file])271 if self.mmproj_url:272 server_args.extend(["--mmproj-url", self.mmproj_url])273 if self.no_mmproj:274 server_args.append("--no-mmproj")275 if self.media_path:276 server_args.extend(["--media-path", self.media_path])277 if self.sleep_idle_seconds is not None:278 server_args.extend(["--sleep-idle-seconds", self.sleep_idle_seconds])279 if self.cache_ram is not None:280 server_args.extend(["--cache-ram", self.cache_ram])281 if self.no_cache_idle_slots:282 server_args.append("--no-cache-idle-slots")283 if self.ui_mcp_proxy:284 server_args.append("--ui-mcp-proxy")285 if self.server_tools:286 server_args.extend(["--tools", self.server_tools])287 if self.server_tools_runtime:288 server_args.extend(["--tools-runtime", self.server_tools_runtime])289 if self.mcp_servers_config:290 server_args.extend(["--mcp-servers-config", self.mcp_servers_config])291 if self.mcp_servers_json:292 server_args.extend(["--mcp-servers-json", self.mcp_servers_json])293 if self.backend_sampling:294 server_args.append("--backend_sampling")295 if self.gcp_compat:296 env["AIP_MODE"] = "PREDICTION"297 env["AIP_HTTP_PORT"] = str(self.server_port)298 299 args = [str(arg) for arg in [server_path, *server_args]]300 print(f"tests: starting server with: {' '.join(args)}")301 302 flags = 0303 if "nt" == os.name:304 flags |= subprocess.DETACHED_PROCESS305 flags |= subprocess.CREATE_NEW_PROCESS_GROUP306 flags |= subprocess.CREATE_NO_WINDOW307 308 if self.log_path:309 self._log = open(self.log_path, "w")310 else:311 self._log = sys.stdout312 313 self.process = subprocess.Popen(314 [str(arg) for arg in [server_path, *server_args]],315 creationflags=flags,316 stdout=self._log,317 stderr=self._log if self._log != sys.stdout else sys.stdout,318 env=env,319 )320 server_instances.add(self)321 322 print(f"server pid={self.process.pid}, pytest pid={os.getpid()}")323 324 # wait for server to start325 start_time = time.time()326 last_print_time = start_time327 while time.time() - start_time < timeout_seconds:328 try:329 response = self.make_request("GET", "/health", headers={330 "Authorization": f"Bearer {self.api_key}" if self.api_key else None331 })332 if response.status_code == 200:333 self.ready = True334 return # server is ready335 except Exception as e:336 pass337 # Check if process died338 if self.process.poll() is not None:339 raise RuntimeError(f"Server process died with return code {self.process.returncode}")340 341 if time.time() - last_print_time >= 1.0:342 print(f"Waiting for server to start...")343 last_print_time = time.time()344 time.sleep(0.01)345 raise TimeoutError(f"Server did not start within {timeout_seconds} seconds")346 347 def stop(self) -> None:348 if self.external_server:349 print("[external_server]: Not stopping external server")350 return351 if self in server_instances:352 server_instances.remove(self)353 if self.process:354 print(f"Stopping server with pid={self.process.pid}")355 self.process.terminate()356 try:357 self.process.wait(timeout=5)358 except subprocess.TimeoutExpired:359 print(f"Server pid={self.process.pid} did not terminate in time, killing")360 self.process.kill()361 self.process.wait(timeout=5)362 except Exception as e:363 print(f"Error waiting for server: {e}")364 self.process = None365 if hasattr(self, '_log') and self._log != sys.stdout:366 self._log.close()367 368 def make_request(369 self,370 method: str,371 path: str,372 data: dict | Any | None = None,373 headers: dict | None = None,374 timeout: float | None = DEFAULT_REQUEST_TIMEOUT,375 ) -> ServerResponse:376 url = f"http://{self.server_host}:{self.server_port}{path}"377 parse_body = False378 if method == "GET":379 response = requests.get(url, headers=headers, timeout=timeout)380 parse_body = True381 elif method == "POST":382 response = requests.post(url, headers=headers, json=data, timeout=timeout)383 parse_body = True384 elif method == "DELETE":385 response = requests.delete(url, headers=headers, timeout=timeout)386 parse_body = True387 elif method == "OPTIONS":388 response = requests.options(url, headers=headers, timeout=timeout)389 else:390 raise ValueError(f"Unimplemented method: {method}")391 result = ServerResponse()392 result.headers = dict(response.headers)393 result.status_code = response.status_code394 if parse_body:395 try:396 result.body = response.json()397 except (JSONDecodeError, requests.exceptions.JSONDecodeError):398 result.body = response.text399 else:400 result.body = None401 print("Response from server", json.dumps(result.body, indent=2))402 return result403 404 def make_stream_request(405 self,406 method: str,407 path: str,408 data: dict | None = None,409 headers: dict | None = None,410 ) -> Iterator[dict]:411 url = f"http://{self.server_host}:{self.server_port}{path}"412 if method == "POST":413 response = requests.post(url, headers=headers, json=data, stream=True)414 else:415 raise ValueError(f"Unimplemented method: {method}")416 if response.status_code != 200:417 raise ServerError(response.status_code, response.json())418 for line_bytes in response.iter_lines():419 line = line_bytes.decode("utf-8")420 if '[DONE]' in line:421 break422 elif line.startswith('data: '):423 data = json.loads(line[6:])424 print("Partial response from server", json.dumps(data, indent=2))425 yield data426 427 def make_any_request(428 self,429 method: str,430 path: str,431 data: dict | None = None,432 headers: dict | None = None,433 timeout: float | None = DEFAULT_REQUEST_TIMEOUT,434 ) -> dict:435 stream = data.get('stream', False)436 if stream:437 content: list[str] = []438 reasoning_content: list[str] = []439 tool_calls: list[dict] = []440 finish_reason: Optional[str] = None441 442 content_parts = 0443 reasoning_content_parts = 0444 tool_call_parts = 0445 arguments_parts = 0446 447 for chunk in self.make_stream_request(method, path, data, headers):448 if chunk['choices']:449 assert len(chunk['choices']) == 1, f'Expected 1 choice, got {len(chunk["choices"])}'450 choice = chunk['choices'][0]451 if choice['delta'].get('content') is not None:452 assert len(choice['delta']['content']) > 0, f'Expected non empty content delta!'453 content.append(choice['delta']['content'])454 content_parts += 1455 if choice['delta'].get('reasoning_content') is not None:456 assert len(choice['delta']['reasoning_content']) > 0, f'Expected non empty reasoning_content delta!'457 reasoning_content.append(choice['delta']['reasoning_content'])458 reasoning_content_parts += 1459 if choice['delta'].get('finish_reason') is not None:460 finish_reason = choice['delta']['finish_reason']461 for tc in choice['delta'].get('tool_calls', []):462 if 'function' not in tc:463 raise ValueError(f"Expected function type, got {tc['type']}")464 if tc['index'] >= len(tool_calls):465 assert 'id' in tc466 assert tc.get('type') == 'function'467 assert 'function' in tc and 'name' in tc['function'] and len(tc['function']['name']) > 0, \468 f"Expected function call with name, got {tc.get('function')}"469 tool_calls.append(dict(470 id="",471 type="function",472 function=dict(473 name="",474 arguments="",475 )476 ))477 tool_call = tool_calls[tc['index']]478 if tc.get('id') is not None:479 tool_call['id'] = tc['id']480 fct = tc['function']481 assert 'id' not in fct, f"Function call should not have id: {fct}"482 if fct.get('name') is not None:483 tool_call['function']['name'] = tool_call['function'].get('name', '') + fct['name']484 if fct.get('arguments') is not None:485 tool_call['function']['arguments'] += fct['arguments']486 arguments_parts += 1487 tool_call_parts += 1488 else:489 # When `include_usage` is True (the default), we expect the last chunk of the stream490 # immediately preceding the `data: [DONE]` message to contain a `choices` field with an empty array491 # and a `usage` field containing the usage statistics (n.b., llama-server also returns `timings` in492 # the last chunk)493 assert 'usage' in chunk, f"Expected finish_reason in chunk: {chunk}"494 assert 'timings' in chunk, f"Expected finish_reason in chunk: {chunk}"495 print(f'Streamed response had {content_parts} content parts, {reasoning_content_parts} reasoning_content parts, {tool_call_parts} tool call parts incl. {arguments_parts} arguments parts')496 result = dict(497 choices=[498 dict(499 index=0,500 finish_reason=finish_reason,501 message=dict(502 role='assistant',503 content=''.join(content) if content else None,504 reasoning_content=''.join(reasoning_content) if reasoning_content else None,505 tool_calls=tool_calls if tool_calls else None,506 ),507 )508 ],509 )510 print("Final response from server", json.dumps(result, indent=2))511 return result512 else:513 response = self.make_request(method, path, data, headers, timeout=timeout)514 assert response.status_code == 200, f"Server returned error: {response.status_code}"515 return response.body516 517 518 519server_instances: Set[ServerProcess] = set()520 521 522class ServerPreset:523 @staticmethod524 def load_all() -> None:525 """ Load all server presets to ensure model files are cached. """526 servers: List[ServerProcess] = [527 method()528 for name, method in ServerPreset.__dict__.items()529 if callable(method) and name != "load_all"530 ]531 for server in servers:532 server.offline = False533 server.start()534 server.stop()535 536 @staticmethod537 def tinyllama2() -> ServerProcess:538 server = ServerProcess()539 server.offline = True # will be downloaded by load_all()540 server.model_hf_repo = "ggml-org/test-model-stories260K"541 server.model_hf_file = None542 server.model_alias = "tinyllama-2"543 server.n_ctx = 512544 server.n_batch = 32545 server.n_slots = 2546 server.n_predict = 64547 server.seed = 42548 return server549 550 @staticmethod551 def bert_bge_small() -> ServerProcess:552 server = ServerProcess()553 server.offline = True # will be downloaded by load_all()554 server.model_hf_repo = "ggml-org/models"555 server.model_hf_file = "bert-bge-small/ggml-model-f16.gguf"556 server.model_alias = "bert-bge-small"557 server.n_ctx = 512558 server.n_batch = 128559 server.n_ubatch = 128560 server.n_slots = 2561 server.seed = 42562 server.server_embeddings = True563 return server564 565 @staticmethod566 def bert_bge_small_with_fa() -> ServerProcess:567 server = ServerProcess()568 server.offline = True # will be downloaded by load_all()569 server.model_hf_repo = "ggml-org/models"570 server.model_hf_file = "bert-bge-small/ggml-model-f16.gguf"571 server.model_alias = "bert-bge-small"572 server.n_ctx = 1024573 server.n_batch = 300574 server.n_ubatch = 300575 server.n_slots = 2576 server.fa = "on"577 server.seed = 42578 server.server_embeddings = True579 return server580 581 @staticmethod582 def tinyllama_infill() -> ServerProcess:583 server = ServerProcess()584 server.offline = True # will be downloaded by load_all()585 server.model_hf_repo = "ggml-org/test-model-stories260K-infill"586 server.model_hf_file = None587 server.model_alias = "tinyllama-infill"588 server.n_ctx = 2048589 server.n_batch = 1024590 server.n_slots = 1591 server.n_predict = 64592 server.temperature = 0.0593 server.seed = 42594 return server595 596 @staticmethod597 def stories15m_moe() -> ServerProcess:598 server = ServerProcess()599 server.offline = True # will be downloaded by load_all()600 server.model_hf_repo = "ggml-org/stories15M_MOE"601 server.model_hf_file = "stories15M_MOE-F16.gguf"602 server.model_alias = "stories15m-moe"603 server.n_ctx = 2048604 server.n_batch = 1024605 server.n_slots = 1606 server.n_predict = 64607 server.temperature = 0.0608 server.seed = 42609 return server610 611 @staticmethod612 def jina_reranker_tiny() -> ServerProcess:613 server = ServerProcess()614 server.offline = True # will be downloaded by load_all()615 server.model_hf_repo = "ggml-org/models"616 server.model_hf_file = "jina-reranker-v1-tiny-en/ggml-model-f16.gguf"617 server.model_alias = "jina-reranker"618 server.n_ctx = 512619 server.n_batch = 512620 server.n_slots = 1621 server.seed = 42622 server.server_reranking = True623 return server624 625 @staticmethod626 def tinygemma3() -> ServerProcess:627 server = ServerProcess()628 server.offline = True # will be downloaded by load_all()629 # mmproj is already provided by HF registry API630 server.model_hf_file = None631 server.model_hf_repo = "ggml-org/tinygemma3-GGUF:Q8_0"632 server.model_alias = "tinygemma3"633 server.n_ctx = 1024634 server.n_batch = 512635 server.n_slots = 2636 server.n_predict = 4637 server.seed = 42638 return server639 640 @staticmethod641 def router() -> ServerProcess:642 server = ServerProcess()643 server.offline = True # will be downloaded by load_all()644 # router server has no models645 server.model_file = None646 server.model_alias = None647 server.model_hf_repo = None648 server.model_hf_file = None649 server.n_ctx = 1024650 server.n_batch = 16651 server.n_slots = 1652 server.n_predict = 16653 server.seed = 42654 return server655 656 657def parallel_function_calls(function_list: List[Tuple[Callable[..., Any], Tuple[Any, ...]]]) -> List[Any]:658 """659 Run multiple functions in parallel and return results in the same order as calls. Equivalent to Promise.all in JS.660 661 Example usage:662 663 results = parallel_function_calls([664 (func1, (arg1, arg2)),665 (func2, (arg3, arg4)),666 ])667 """668 results = [None] * len(function_list)669 exceptions = []670 671 def worker(index, func, args):672 try:673 result = func(*args)674 results[index] = result675 except Exception as e:676 exceptions.append((index, str(e)))677 678 with ThreadPoolExecutor() as executor:679 futures = []680 for i, (func, args) in enumerate(function_list):681 future = executor.submit(worker, i, func, args)682 futures.append(future)683 684 # Wait for all futures to complete685 for future in as_completed(futures):686 pass687 688 # Check if there were any exceptions689 if exceptions:690 print("Exceptions occurred:")691 for index, error in exceptions:692 print(f"Function at index {index}: {error}")693 694 return results695 696 697def match_regex(regex: str, text: str) -> bool:698 return (699 re.compile(700 regex, flags=RegexFlag.IGNORECASE | RegexFlag.MULTILINE | RegexFlag.DOTALL701 ).search(text)702 is not None703 )704 705 706def download_file(url: str, output_file_path: str | None = None) -> str:707 """708 Download a file from a URL to a local path. If the file already exists, it will not be downloaded again.709 710 output_file_path is the local path to save the downloaded file. If not provided, the file will be saved in the root directory.711 712 Returns the local path of the downloaded file.713 """714 file_name = url.split('/').pop()715 output_file = f'./tmp/{file_name}' if output_file_path is None else output_file_path716 if not os.path.exists(output_file):717 print(f"Downloading {url} to {output_file}")718 wget.download(url, out=output_file)719 print(f"Done downloading to {output_file}")720 else:721 print(f"File already exists at {output_file}")722 return output_file723 724 725def is_slow_test_allowed():726 return os.environ.get("SLOW_TESTS") == "1" or os.environ.get("SLOW_TESTS") == "ON"727 