unsloth/studio/backend/tests/test_llama_cpp_no_context_shift.py
Daniel Han-Chen 05184ad15a Fix llama_cpp source-inspection tests for split load_model for PR #5754
Round 15 split LlamaCppBackend.load_model into a thin wrapper that
publishes _loading_model_identifier + _loading_hf_variant under
_serial_load_lock and an inner _load_model_impl_locked body that
actually launches llama-server. The pre-existing source-inspection
regression tests inspected only load_model and broke because the
flag literals and _wait_for_vram_settle call now live in the inner
method:

- tests/test_llama_cpp_no_context_shift.py
  test_no_context_shift_is_in_load_model
  test_flag_sits_inside_the_base_cmd_list
- tests/test_llama_cpp_wait_for_vram_settle.py
  test_load_model_calls_helper_outside_lock_and_uses_last_kill_timestamp

Update both helpers to concatenate the source of load_model AND
_load_model_impl_locked so the assertions still cover the launch
path without weakening their scope to the full module.
2026-05-25 07:15:16 +00:00

148 lines
5.6 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""``--no-context-shift`` launch-flag contract.
When llama-server runs with its default context-shift behavior, the UI
has no way to tell the user that the KV cache has been rotated --
earlier turns silently vanish from the conversation. The Studio
backend always passes ``--no-context-shift`` so the server returns a
clean error instead, and the chat adapter can point the user at the
``Context Length`` input in the settings panel.
This file is a static read of the launch command: we ask
``LlamaCppBackend`` to assemble its ``cmd`` list and assert the flag
is always present. Testing via the real subprocess would require an
actual GGUF on disk, which is out of scope for the fast test suite.
"""
from __future__ import annotations
import inspect
import sys
import types as _types
from pathlib import Path
import pytest
# ---------------------------------------------------------------------------
# Same external-dep stubs as the other llama_cpp tests.
# ---------------------------------------------------------------------------
_BACKEND_DIR = str(Path(__file__).resolve().parent.parent)
if _BACKEND_DIR not in sys.path:
sys.path.insert(0, _BACKEND_DIR)
_loggers_stub = _types.ModuleType("loggers")
_loggers_stub.get_logger = lambda name: __import__("logging").getLogger(name)
sys.modules.setdefault("loggers", _loggers_stub)
_structlog_stub = _types.ModuleType("structlog")
sys.modules.setdefault("structlog", _structlog_stub)
_httpx_stub = _types.ModuleType("httpx")
for _exc in (
"ConnectError",
"TimeoutException",
"ReadTimeout",
"ReadError",
"RemoteProtocolError",
"CloseError",
):
setattr(_httpx_stub, _exc, type(_exc, (Exception,), {}))
_httpx_stub.Timeout = type("T", (), {"__init__": lambda s, *a, **k: None})
_httpx_stub.Client = type(
"C",
(),
{
"__init__": lambda s, **kw: None,
"__enter__": lambda s: s,
"__exit__": lambda s, *a: None,
},
)
sys.modules.setdefault("httpx", _httpx_stub)
from core.inference import llama_cpp as llama_cpp_module
def _load_model_source() -> str:
"""Return the source of ``LlamaCppBackend.load_model`` PLUS the
internal ``_load_model_impl_locked`` body it delegates to.
Studio's diffusion PR split ``load_model`` into a thin wrapper
that publishes ``_loading_model_identifier`` under
``_serial_load_lock`` and an inner ``_load_model_impl_locked``
body that actually spawns llama-server. The launch flags and the
``_wait_for_vram_settle`` call now live in the inner method, so
inspecting only ``load_model`` would miss them. Concatenating the
two sources keeps these source-inspection regression tests
working without weakening the scope (we still only look at the
two load entry points, not the entire module).
"""
parts = [inspect.getsource(llama_cpp_module.LlamaCppBackend.load_model)]
impl = getattr(
llama_cpp_module.LlamaCppBackend, "_load_model_impl_locked", None
)
if impl is not None:
parts.append(inspect.getsource(impl))
return "\n".join(parts)
def test_no_context_shift_is_in_load_model():
"""The flag is part of the static launch-command template.
We check the source of ``load_model`` rather than mocking the whole
call chain (GPU probing, GGUF stat, etc.): the flag is written as
a literal in one place and any regression has to delete it, which
a text search will catch.
"""
assert '"--no-context-shift"' in _load_model_source(), (
"llama-server must be launched with --no-context-shift so the "
"UI can surface a clean 'context full' error instead of silently "
"losing old turns to a KV-cache rotation."
)
def test_flag_sits_inside_the_base_cmd_list():
"""Pin the flag's location so a future refactor can't accidentally
move it into a branch that only fires on some code paths.
We slice from ``cmd = [`` to the first ``]`` at the same indent.
Using ``inspect.getsource`` means the function lives in its own
string and there are no siblings to worry about, so a plain
bracket search would also work -- anchoring on the trailing indent
just keeps the slice from wandering into a later expression if the
opening literal ever grows an in-line comment trailing it.
"""
source = _load_model_source()
start = source.find("cmd = [")
assert start >= 0, "could not find the base cmd = [...] block"
# Find the first line containing only ``]`` (possibly indented).
# Works for any indentation style the formatter picks.
rest = source[start:]
end_rel = -1
for line_start, line in _iter_lines_with_offset(rest):
if line_start == 0:
# Skip the opening ``cmd = [`` line itself.
continue
if line.strip() == "]":
end_rel = line_start
break
assert end_rel > 0, "could not find end of cmd = [...] block"
block = rest[:end_rel]
assert '"--no-context-shift"' in block, (
"--no-context-shift must be in the base cmd list, not in a "
"conditional branch -- otherwise some code paths would still "
"run with silent context shift enabled."
)
# Also pin that it is next to -c / --ctx so the grouping makes sense.
assert '"-c"' in block
assert '"--flash-attn"' in block
def _iter_lines_with_offset(text: str):
"""Yield (offset, line) pairs over ``text`` without losing offsets."""
offset = 0
for line in text.splitlines(keepends = True):
yield offset, line
offset += len(line)