Files
cleveragents-core/features/steps/lsp_server_stub_steps.py
HAL9000 22c3cddf08
CI / lint (pull_request) Successful in 39s
CI / quality (pull_request) Successful in 51s
CI / typecheck (pull_request) Successful in 1m9s
CI / build (pull_request) Successful in 32s
CI / helm (pull_request) Successful in 37s
CI / security (pull_request) Successful in 1m34s
CI / push-validation (pull_request) Successful in 31s
CI / unit_tests (pull_request) Successful in 4m38s
CI / docker (pull_request) Successful in 1m54s
CI / integration_tests (pull_request) Successful in 8m35s
CI / coverage (pull_request) Successful in 9m1s
CI / status-check (pull_request) Successful in 3s
test(lsp): cover the select-based read-body timeout branch via os.pipe()
The DoS mitigation added in db389a730 wraps each message-body read in
``select()`` so a stalled client cannot pin the server forever. The
existing transport tests route through MockLspTransport, whose
``BytesIO`` raises ``UnsupportedOperation`` on ``fileno()`` -- so
``_read_body_with_timeout`` always falls through the BytesIO fast
path and the actual ``select()``-based mitigation code (the part
that runs in production) is never executed. That left ~19 new lines
uncovered, dragging total coverage below the 96.5% floor and
failing the coverage gate.

Two new scenarios drive the helper through an ``os.pipe()`` whose
read fd satisfies ``fileno()``, so ``use_select`` is True and the
real DoS-protection path runs:

* ``timeout=0.0`` makes the deadline already past on the first
  iteration, exercising the ``if timeout <= 0:`` early-exit warning.
* ``timeout=0.05`` lets ``select.select()`` run and time out with no
  ready descriptors, exercising the ``if not ready:`` warning.

Both paths log ``lsp.transport.read_timeout`` and return ``None``,
matching the production behaviour the helper was added to provide.

ISSUES CLOSED: #5566
2026-06-05 17:25:35 -04:00

880 lines
33 KiB
Python

"""Step definitions for LSP server stub BDD tests.
Covers the JSON-RPC transport, protocol handshake, error handling,
A2A facade wiring, and CLI serve command for the LSP server stub.
"""
from __future__ import annotations
import contextlib
import io
import json
import os
from collections.abc import Generator
from typing import Any
import structlog
from behave import given, then, when
from behave.runner import Context
from typer.testing import CliRunner
from cleveragents.a2a.facade import A2aLocalFacade
from cleveragents.cli.commands.lsp import app as lsp_app
from cleveragents.lsp import server as _server_module
from cleveragents.lsp.server import (
MAX_CONTENT_LENGTH,
MAX_HEADER_LINE_LENGTH,
MAX_HEADER_LINES,
LspServer,
)
from features.mocks.lsp_transport_mock import MockLspTransport
@contextlib.contextmanager
def _capture_structlogs() -> Generator[list[dict[str, Any]]]:
"""Capture structlog entries emitted by the LSP server module.
``structlog.testing.capture_logs()`` does not work when the global
configuration has ``cache_logger_on_first_use=True`` because the
module-level ``logger`` in ``cleveragents.lsp.server`` caches its
processor chain on first use and never picks up the replacement
processors installed by ``capture_logs()``.
This helper works around the problem by:
1. Saving and temporarily replacing the module-level ``logger``
with a fresh proxy obtained *after* reconfiguring structlog
with ``cache_logger_on_first_use=False``.
2. Restoring both the structlog configuration and the original
module ``logger`` when the context manager exits.
"""
cap = structlog.testing.LogCapture()
old_config = structlog.get_config()
old_logger = _server_module.logger
structlog.configure(
processors=[cap],
wrapper_class=structlog.stdlib.BoundLogger,
logger_factory=structlog.PrintLoggerFactory(),
cache_logger_on_first_use=False,
)
# Replace the module-level logger so calls inside ``server.py``
# go through the new (uncached) configuration.
_server_module.logger = structlog.get_logger(_server_module.__name__) # type: ignore[assignment]
try:
yield cap.entries
finally:
_server_module.logger = old_logger
structlog.configure(**old_config)
_cli_runner = CliRunner()
# ---------------------------------------------------------------------------
# Background
# ---------------------------------------------------------------------------
@given("a mock LSP transport")
def step_mock_transport(context: Context) -> None:
context.lsp_transport = MockLspTransport()
context.lsp_server: LspServer | None = None
context.lsp_responses: list[dict[str, Any]] = []
context.lsp_exit_code: int | None = None
context.lsp_facade: A2aLocalFacade | None = None
context.lsp_error: Exception | None = None
# ---------------------------------------------------------------------------
# Given — messages
# ---------------------------------------------------------------------------
@given("an initialize request with id {req_id:d}")
def step_initialize_request(context: Context, req_id: int) -> None:
context.lsp_transport.send_request(
"initialize",
{"capabilities": {}},
request_id=req_id,
)
@given('an initialize request with id {req_id:d} and clientInfo name "{name}"')
def step_initialize_with_client_info(context: Context, req_id: int, name: str) -> None:
context.lsp_transport.send_request(
"initialize",
{"capabilities": {}, "clientInfo": {"name": name}},
request_id=req_id,
)
@given("a shutdown request with id {req_id:d}")
def step_shutdown_request(context: Context, req_id: int) -> None:
context.lsp_transport.send_request("shutdown", request_id=req_id)
@given("an exit notification")
def step_exit_notification(context: Context) -> None:
context.lsp_transport.send_notification("exit")
@given('a request for method "{method}" with id {req_id:d}')
def step_request_method(context: Context, method: str, req_id: int) -> None:
context.lsp_transport.send_request(method, {}, request_id=req_id)
@given('a notification for method "{method}"')
def step_notification_method(context: Context, method: str) -> None:
context.lsp_transport.send_notification(method)
@given("raw bytes that are not valid JSON-RPC")
def step_raw_invalid_bytes(context: Context) -> None:
bad_payload = b"this is not json at all"
framed = f"Content-Length: {len(bad_payload)}\r\n\r\n".encode() + bad_payload
context.lsp_transport.send_raw(framed)
@given("a message with missing jsonrpc version and id {req_id:d}")
def step_missing_jsonrpc_version(context: Context, req_id: int) -> None:
msg = {"id": req_id, "method": "test"}
payload = json.dumps(msg).encode("utf-8")
framed = f"Content-Length: {len(payload)}\r\n\r\n".encode() + payload
context.lsp_transport.send_raw(framed)
@given("a message with missing method and id {req_id:d}")
def step_missing_method(context: Context, req_id: int) -> None:
msg = {"jsonrpc": "2.0", "id": req_id}
payload = json.dumps(msg).encode("utf-8")
framed = f"Content-Length: {len(payload)}\r\n\r\n".encode() + payload
context.lsp_transport.send_raw(framed)
@given("an LSP server with A2A facade")
def step_server_with_facade(context: Context) -> None:
context.lsp_facade = A2aLocalFacade()
# ---------------------------------------------------------------------------
# When — server execution
# ---------------------------------------------------------------------------
@when("the LSP server processes all messages")
def step_process_messages(context: Context) -> None:
facade = getattr(context, "lsp_facade", None)
server = LspServer(
input_stream=context.lsp_transport.input_stream,
output_stream=context.lsp_transport.output_stream,
facade=facade,
)
context.lsp_server = server
with _capture_structlogs() as captured:
context.lsp_exit_code = server.run()
context.lsp_captured_logs = captured
context.lsp_responses = context.lsp_transport.read_responses()
@when("the LSP server processes all messages with empty input")
def step_process_empty(context: Context) -> None:
empty_input = io.BytesIO(b"")
output_buf = io.BytesIO()
server = LspServer(input_stream=empty_input, output_stream=output_buf)
context.lsp_server = server
context.lsp_exit_code = server.run()
context.lsp_responses = []
@when("the LSP server processes all messages with a broken output stream")
def step_process_broken_output(context: Context) -> None:
"""Run the server with an output stream that raises OSError on write."""
class _BrokenOutput:
"""Mock output stream that raises on any write."""
def write(self, _data: bytes) -> int:
raise OSError("broken pipe")
def flush(self) -> None:
raise OSError("broken pipe")
server = LspServer(
input_stream=context.lsp_transport.input_stream,
output_stream=_BrokenOutput(), # type: ignore[arg-type]
)
context.lsp_server = server
with _capture_structlogs() as captured:
context.lsp_exit_code = server.run()
context.lsp_captured_logs = captured
context.lsp_responses = []
@when("I try to create an LspServer with invalid input stream")
def step_invalid_input_stream(context: Context) -> None:
try:
LspServer(input_stream="not a stream") # type: ignore[arg-type]
context.lsp_error = None
except TypeError as exc:
context.lsp_error = exc
@when("I try to create an LspServer with invalid output stream")
def step_invalid_output_stream(context: Context) -> None:
try:
LspServer(output_stream=42) # type: ignore[arg-type]
context.lsp_error = None
except TypeError as exc:
context.lsp_error = exc
# ---------------------------------------------------------------------------
# Then — response assertions
# ---------------------------------------------------------------------------
def _find_response(context: Context, req_id: int) -> dict[str, Any]:
for resp in context.lsp_responses:
if resp.get("id") == req_id:
return resp
raise AssertionError(
f"No response with id={req_id} found in {context.lsp_responses}"
)
def _find_response_any(
context: Context,
req_id: int | str,
) -> dict[str, Any]:
"""Find a response by id, supporting both int and str ids."""
for resp in context.lsp_responses:
if resp.get("id") == req_id:
return resp
raise AssertionError(
f"No response with id={req_id!r} found in {context.lsp_responses}"
)
@then("the response for id {req_id:d} should have result")
def step_response_has_result(context: Context, req_id: int) -> None:
resp = _find_response(context, req_id)
assert "result" in resp, f"Response {resp} missing 'result'"
assert resp["result"] is not None, f"Response result is None: {resp}"
@then("the response for id {req_id:d} should have null result")
def step_response_null_result(context: Context, req_id: int) -> None:
resp = _find_response(context, req_id)
assert "result" in resp, f"Response {resp} missing 'result'"
assert resp["result"] is None, f"Expected null result, got {resp['result']}"
@then('the result for id {req_id:d} should contain serverInfo with name "{name}"')
def step_result_server_name(context: Context, req_id: int, name: str) -> None:
resp = _find_response(context, req_id)
result = resp.get("result", {})
server_info = result.get("serverInfo", {})
assert server_info.get("name") == name, (
f"Expected serverInfo.name={name}, got {server_info}"
)
@then('the result for id {req_id:d} should contain serverInfo with version "{version}"')
def step_result_server_version(context: Context, req_id: int, version: str) -> None:
resp = _find_response(context, req_id)
result = resp.get("result", {})
server_info = result.get("serverInfo", {})
assert server_info.get("version") == version, (
f"Expected version={version}, got {server_info}"
)
@then('the result for id {req_id:d} should contain capabilities key "{key}"')
def step_result_capabilities_key(context: Context, req_id: int, key: str) -> None:
resp = _find_response(context, req_id)
result = resp.get("result", {})
caps = result.get("capabilities", {})
assert key in caps, f"Capabilities missing key '{key}': {caps}"
@then("the response for id {req_id:d} should have error code {code:d}")
def step_response_error_code(context: Context, req_id: int, code: int) -> None:
resp = _find_response(context, req_id)
assert "error" in resp, f"Response {resp} missing 'error'"
assert resp["error"]["code"] == code, (
f"Expected error code {code}, got {resp['error']['code']}"
)
@then('the error message for id {req_id:d} should contain "{text}"')
def step_error_message_for_id(context: Context, req_id: int, text: str) -> None:
resp = _find_response(context, req_id)
error = resp.get("error")
assert error is not None, f"Response for id={req_id} has no error: {resp}"
msg = error.get("message", "")
assert text in msg, (
f"Error message for id={req_id} does not contain '{text}': {msg}"
)
@then("a parse error response should have been sent")
def step_parse_error_sent(context: Context) -> None:
errors = [
r
for r in context.lsp_responses
if "error" in r and r["error"].get("code") == -32700
]
assert errors, f"No parse error (-32700) response found in {context.lsp_responses}"
@then("there should be no error responses")
def step_no_error_responses(context: Context) -> None:
errors = [r for r in context.lsp_responses if "error" in r]
assert not errors, f"Unexpected error responses: {errors}"
@then("the server exit code should be {code:d}")
def step_exit_code(context: Context, code: int) -> None:
assert context.lsp_exit_code == code, (
f"Expected exit code {code}, got {context.lsp_exit_code}"
)
@then("the server should have stopped")
def step_server_stopped(context: Context) -> None:
assert context.lsp_server is not None, "No server was created"
assert context.lsp_server.is_running is False, "Server is still running"
@then('a transport warning should have been logged for "{event}"')
def step_transport_warning_logged(context: Context, event: str) -> None:
captured = getattr(context, "lsp_captured_logs", [])
matches = [
e
for e in captured
if e.get("event") == event and e.get("log_level", "") == "warning"
]
assert matches, f"No warning log with event={event!r} found in {captured}"
@then("the server should have logged its PID")
def step_server_logged_pid(context: Context) -> None:
assert context.lsp_server is not None, "No server was created"
assert context.lsp_exit_code is not None, "Server did not complete"
captured = getattr(context, "lsp_captured_logs", [])
pid_entries = [
e for e in captured if "pid" in e and e.get("event") == "lsp.server.starting"
]
assert pid_entries, (
f"No log entry with PID found for 'lsp.server.starting' in {captured}"
)
assert isinstance(pid_entries[0]["pid"], int), "PID should be an integer"
@then("an InvalidRequest error response should have been sent")
def step_invalid_request_sent(context: Context) -> None:
errors = [
r
for r in context.lsp_responses
if "error" in r and r["error"].get("code") == -32600
]
assert errors, (
f"No InvalidRequest (-32600) response found in {context.lsp_responses}"
)
@then("the error response for id {req_id:d} should not contain a data field")
def step_error_no_data(context: Context, req_id: int) -> None:
resp = _find_response(context, req_id)
assert "error" in resp, f"Response {resp} missing 'error'"
assert "data" not in resp["error"], (
f"Error response unexpectedly contains 'data': {resp['error']}"
)
@then("the server should hold the provided A2A facade")
def step_server_holds_facade(context: Context) -> None:
assert context.lsp_server is not None, "No server was created"
assert context.lsp_facade is not None, "No facade was provided"
assert context.lsp_server.facade is context.lsp_facade, (
"Server does not hold the provided A2A facade instance"
)
@then("the server facade should be lazily created")
def step_server_facade_lazy(context: Context) -> None:
assert context.lsp_server is not None, "No server was created"
# Access the public property — this triggers lazy creation
facade = context.lsp_server.facade
assert facade is not None, "Facade was not lazily created"
from cleveragents.a2a.facade import A2aLocalFacade
assert isinstance(facade, A2aLocalFacade), (
f"Expected A2aLocalFacade, got {type(facade).__name__}"
)
@then("a TypeError should be raised for input stream")
def step_type_error_input(context: Context) -> None:
assert context.lsp_error is not None, "Expected TypeError"
assert isinstance(context.lsp_error, TypeError), (
f"Expected TypeError, got {type(context.lsp_error).__name__}"
)
assert "input_stream" in str(context.lsp_error)
@then("a TypeError should be raised for output stream")
def step_type_error_output(context: Context) -> None:
assert context.lsp_error is not None, "Expected TypeError"
assert isinstance(context.lsp_error, TypeError), (
f"Expected TypeError, got {type(context.lsp_error).__name__}"
)
assert "output_stream" in str(context.lsp_error)
# ---------------------------------------------------------------------------
# Transport edge case steps
# ---------------------------------------------------------------------------
@given("a request with non-dict params and id {req_id:d}")
def step_non_dict_params(context: Context, req_id: int) -> None:
msg: dict[str, Any] = {
"jsonrpc": "2.0",
"id": req_id,
"method": "textDocument/completion",
"params": "not-a-dict",
}
payload = json.dumps(msg).encode("utf-8")
framed = f"Content-Length: {len(payload)}\r\n\r\n".encode() + payload
context.lsp_transport.send_raw(framed)
@given("a message with a JSON array body")
def step_json_array_body(context: Context) -> None:
# Send a valid JSON array instead of a JSON object so that
# _handle_message's isinstance(msg, dict) guard is exercised.
payload = json.dumps([1, "not", "a", "dict"]).encode("utf-8")
framed = f"Content-Length: {len(payload)}\r\n\r\n".encode() + payload
context.lsp_transport.send_raw(framed)
@given("a message with non-integer content length")
def step_non_integer_content_length(context: Context) -> None:
raw = b"Content-Length: abc\r\n\r\n"
context.lsp_transport.send_raw(raw)
@given("a message with initialize request without params key and id {req_id:d}")
def step_initialize_no_params(context: Context, req_id: int) -> None:
msg: dict[str, Any] = {
"jsonrpc": "2.0",
"id": req_id,
"method": "initialize",
}
payload = json.dumps(msg).encode("utf-8")
framed = f"Content-Length: {len(payload)}\r\n\r\n".encode() + payload
context.lsp_transport.send_raw(framed)
@given("a message with negative content length")
def step_negative_content_length(context: Context) -> None:
raw = b"Content-Length: -5\r\n\r\n"
context.lsp_transport.send_raw(raw)
@given("a message with incomplete body")
def step_incomplete_body(context: Context) -> None:
# Declare 100 bytes but only provide 5
raw = b"Content-Length: 100\r\n\r\nhello"
context.lsp_transport.send_raw(raw)
@given("a message with missing content length header")
def step_missing_content_length(context: Context) -> None:
raw = b"Content-Type: application/json\r\n\r\n{}"
context.lsp_transport.send_raw(raw)
@given("a message with content length exceeding the maximum")
def step_oversized_content_length(context: Context) -> None:
# Declare a Content-Length above the limit. The *actual* body is
# ``MAX_CONTENT_LENGTH`` bytes — matching the server's capped
# discard — so ``_discard_body`` consumes the entire body and
# subsequent messages (the exit notification) are not corrupted.
declared_length = MAX_CONTENT_LENGTH + 64
body = b"X" * MAX_CONTENT_LENGTH
raw = f"Content-Length: {declared_length}\r\n\r\n".encode() + body
context.lsp_transport.send_raw(raw)
@given("a message with content length zero")
def step_zero_content_length(context: Context) -> None:
raw = b"Content-Length: 0\r\n\r\n"
context.lsp_transport.send_raw(raw)
@given("a message with duplicate content-length headers and id {req_id:d}")
def step_duplicate_content_length(context: Context, req_id: int) -> None:
# Build an initialize request. The first Content-Length is wrong
# (too small), but the server takes the *last* value so the
# message is parsed correctly.
msg: dict[str, Any] = {
"jsonrpc": "2.0",
"id": req_id,
"method": "initialize",
"params": {"capabilities": {}},
}
payload = json.dumps(msg).encode("utf-8")
header = (f"Content-Length: 1\r\nContent-Length: {len(payload)}\r\n\r\n").encode()
context.lsp_transport.send_raw(header + payload)
@given("a message with deeply nested JSON")
def step_deeply_nested_json(context: Context) -> None:
"""Build a Content-Length framed message whose body is a deeply nested
JSON array (10 000 levels). When ``json.loads`` attempts to parse
this payload it triggers a ``RecursionError``, exercising the guard
in ``_read_message``.
"""
depth = 10000
payload = ("[" * depth + "]" * depth).encode("utf-8")
framed = f"Content-Length: {len(payload)}\r\n\r\n".encode() + payload
context.lsp_transport.send_raw(framed)
@given("an oversized message with a short body that triggers EOF during discard")
def step_oversized_short_body_eof(context: Context) -> None:
"""Build a message whose declared CL exceeds the limit but whose actual
body is much shorter than the discard cap. This forces
``_discard_body`` to hit EOF mid-read (the early-exit ``break``
inside ``LspServer._discard_body``'s read loop —
``cleveragents.lsp.server.LspServer._discard_body``), covering
the early-exit path.
"""
declared_length = MAX_CONTENT_LENGTH + 64
# Only provide a tiny body — far less than the capped discard amount
body = b"X" * 128
raw = f"Content-Length: {declared_length}\r\n\r\n".encode() + body
context.lsp_transport.send_raw(raw)
@given("a message with too many header lines")
def step_too_many_header_lines(context: Context) -> None:
lines = [f"X-Spam-{i}: value\r\n" for i in range(MAX_HEADER_LINES + 1)]
raw = "".join(lines).encode()
context.lsp_transport.send_raw(raw)
@given("a message with a header line exceeding the maximum line length")
def step_oversized_header_line(context: Context) -> None:
"""Build a single header line that exceeds ``MAX_HEADER_LINE_LENGTH``.
When ``readline(MAX_HEADER_LINE_LENGTH)`` is used the oversized
line is returned in multiple chunks, each consuming one slot in
the ``MAX_HEADER_LINES`` budget. Eventually the header-flood
guard fires, exercising the ``readline()`` per-line size cap
path described in ``cleveragents.lsp.server.LspServer._read_message``.
"""
# A single line without a newline that far exceeds the per-line
# limit. ``readline(MAX_HEADER_LINE_LENGTH)`` will return it in
# chunks of MAX_HEADER_LINE_LENGTH bytes until MAX_HEADER_LINES
# is exceeded.
oversized_line = b"X-Huge: " + b"A" * (
MAX_HEADER_LINE_LENGTH * (MAX_HEADER_LINES + 2)
)
context.lsp_transport.send_raw(oversized_line)
# ---------------------------------------------------------------------------
# Stream desync recovery steps
# ---------------------------------------------------------------------------
@given(
"an oversized message with real body followed by"
" an initialize request with id {req_id:d}"
)
def step_oversized_with_body_then_init(context: Context, req_id: int) -> None:
"""Build an oversized message with actual body bytes and append a valid init.
The declared Content-Length exceeds the limit so the server rejects the
message. The *actual* body is ``MAX_CONTENT_LENGTH`` bytes — matching
the server's capped discard — so that ``_discard_body`` consumes
exactly the available body and the stream stays aligned for the next
message.
"""
declared_length = MAX_CONTENT_LENGTH + 64
# Actual body matches the discard cap so recovery succeeds.
body = b"X" * MAX_CONTENT_LENGTH
header = f"Content-Length: {declared_length}\r\n\r\n".encode()
raw = header + body
# Append the valid initialize request after the oversized body.
init_msg: dict[str, Any] = {
"jsonrpc": "2.0",
"id": req_id,
"method": "initialize",
"params": {"capabilities": {}},
}
init_payload = json.dumps(init_msg).encode("utf-8")
init_framed = f"Content-Length: {len(init_payload)}\r\n\r\n".encode() + init_payload
context.lsp_transport.send_raw(raw + init_framed)
@given(
"a message with non-integer content length followed by"
" an initialize request with id {req_id:d}"
)
def step_invalid_cl_then_init(context: Context, req_id: int) -> None:
"""Build a message with invalid CL header followed by valid init."""
# Invalid CL header + additional header + blank line (body-less)
bad_raw = b"Content-Length: abc\r\nX-Extra: foo\r\n\r\n"
init_msg: dict[str, Any] = {
"jsonrpc": "2.0",
"id": req_id,
"method": "initialize",
"params": {"capabilities": {}},
}
init_payload = json.dumps(init_msg).encode("utf-8")
init_framed = f"Content-Length: {len(init_payload)}\r\n\r\n".encode() + init_payload
context.lsp_transport.send_raw(bad_raw + init_framed)
@given(
"a message with too many header lines followed by"
" an initialize request with id {req_id:d}"
)
def step_too_many_headers_then_init(context: Context, req_id: int) -> None:
"""Build a message with excess headers + blank terminator, then valid init."""
lines = [f"X-Spam-{i}: value\r\n" for i in range(MAX_HEADER_LINES + 1)]
# Append the blank line terminator so _drain_headers can reach it.
lines.append("\r\n")
bad_raw = "".join(lines).encode()
init_msg: dict[str, Any] = {
"jsonrpc": "2.0",
"id": req_id,
"method": "initialize",
"params": {"capabilities": {}},
}
init_payload = json.dumps(init_msg).encode("utf-8")
init_framed = f"Content-Length: {len(init_payload)}\r\n\r\n".encode() + init_payload
context.lsp_transport.send_raw(bad_raw + init_framed)
# ---------------------------------------------------------------------------
# CLI serve command steps
# ---------------------------------------------------------------------------
@when('I run lsp CLI serve with log level "{level}"')
def step_run_lsp_serve_bad_level(context: Context, level: str) -> None:
result = _cli_runner.invoke(lsp_app, ["serve", "--log-level", level])
context.lsp_cli_result = result
@when("I run lsp CLI serve with piped empty stdin")
def step_run_lsp_serve_empty_stdin(context: Context) -> None:
result = _cli_runner.invoke(lsp_app, ["serve"], input="")
context.lsp_cli_result = result
@when('I run lsp CLI serve with log level "{level}" and piped empty stdin')
def step_run_lsp_serve_level_empty_stdin(context: Context, level: str) -> None:
result = _cli_runner.invoke(lsp_app, ["serve", "--log-level", level], input="")
context.lsp_cli_result = result
@then("the lsp serve CLI should have failed")
def step_lsp_serve_failed(context: Context) -> None:
result = context.lsp_cli_result
assert result is not None, "No CLI result captured"
assert result.exit_code != 0, (
f"Expected non-zero exit code, got {result.exit_code}. Output:\n{result.output}"
)
@then("the lsp serve CLI should have exited")
def step_lsp_serve_exited(context: Context) -> None:
result = context.lsp_cli_result
assert result is not None, "No CLI result captured"
# The server exits with code 0 (clean shutdown) or 1 (no shutdown sent).
assert result.exit_code in (0, 1), (
f"Unexpected exit code {result.exit_code}. Output:\n{result.output}"
)
@then("the lsp serve CLI should have exited with code {code:d}")
def step_lsp_serve_exited_with_code(context: Context, code: int) -> None:
result = context.lsp_cli_result
assert result is not None, "No CLI result captured"
assert result.exit_code == code, (
f"Expected exit code {code}, got {result.exit_code}. Output:\n{result.output}"
)
@then('the lsp serve CLI output should contain "{text}"')
def step_lsp_serve_output_contains(context: Context, text: str) -> None:
# NOTE: The serve command writes its startup banner to stderr via
# ``_serve_console = Console(stderr=True)``. This assertion checks
# ``result.output`` (stdout) which works because ``CliRunner()``
# defaults to ``mix_stderr=True``, merging stderr into stdout.
result = context.lsp_cli_result
assert result is not None, "No CLI result captured"
assert text in result.output, f"Expected '{text}' in output, got:\n{result.output}"
# ---------------------------------------------------------------------------
# String request id steps (JSON-RPC 2.0 allows string ids)
# ---------------------------------------------------------------------------
@given('an initialize request with string id "{req_id}"')
def step_initialize_string_id(context: Context, req_id: str) -> None:
context.lsp_transport.send_request(
"initialize",
{"capabilities": {}},
request_id=req_id,
)
@then('the response for string id "{req_id}" should have result')
def step_response_string_id_has_result(context: Context, req_id: str) -> None:
resp = _find_response_any(context, req_id)
assert "result" in resp, f"Response {resp} missing 'result'"
assert resp["result"] is not None, f"Response result is None: {resp}"
@then(
'the result for string id "{req_id}" should contain serverInfo with name "{name}"'
)
def step_result_string_id_server_name(
context: Context,
req_id: str,
name: str,
) -> None:
resp = _find_response_any(context, req_id)
result = resp.get("result", {})
server_info = result.get("serverInfo", {})
assert server_info.get("name") == name, (
f"Expected serverInfo.name={name}, got {server_info}"
)
# ---------------------------------------------------------------------------
# Per-message read timeout (DoS protection) steps
# ---------------------------------------------------------------------------
@when("the LSP server processes all messages with read timeout {timeout:f}")
def step_process_with_read_timeout(context: Context, timeout: float) -> None:
"""Run the LSP server with a custom read timeout using the mock transport."""
facade = getattr(context, "lsp_facade", None)
server = LspServer(
input_stream=context.lsp_transport.input_stream,
output_stream=context.lsp_transport.output_stream,
facade=facade,
read_timeout=timeout,
)
context.lsp_server = server
with _capture_structlogs() as captured:
context.lsp_exit_code = server.run()
context.lsp_captured_logs = captured
context.lsp_responses = context.lsp_transport.read_responses()
@when("I check the MESSAGE_READ_TIMEOUT constant")
def step_check_message_read_timeout(context: Context) -> None:
"""Store the MESSAGE_READ_TIMEOUT value for assertion."""
from cleveragents.lsp.server import MESSAGE_READ_TIMEOUT
context.lsp_message_read_timeout = MESSAGE_READ_TIMEOUT
@then("the MESSAGE_READ_TIMEOUT should be {value:d}")
def step_message_read_timeout_value(context: Context, value: int) -> None:
"""Assert the MESSAGE_READ_TIMEOUT constant has the expected value."""
assert context.lsp_message_read_timeout == value, (
f"Expected MESSAGE_READ_TIMEOUT={value}, got {context.lsp_message_read_timeout}"
)
@then("MESSAGE_READ_TIMEOUT should be importable from lsp server module")
def step_message_read_timeout_importable(context: Context) -> None:
"""Assert MESSAGE_READ_TIMEOUT is in __all__ and importable."""
from cleveragents.lsp import server as _lsp_server_module
assert "MESSAGE_READ_TIMEOUT" in _lsp_server_module.__all__, (
f"MESSAGE_READ_TIMEOUT not found in __all__: {_lsp_server_module.__all__}"
)
assert hasattr(_lsp_server_module, "MESSAGE_READ_TIMEOUT"), (
"MESSAGE_READ_TIMEOUT attribute not found in lsp.server module"
)
# ---------------------------------------------------------------------------
# Real-pipe timeout coverage for ``_read_body_with_timeout``
# ---------------------------------------------------------------------------
# The MockLspTransport uses ``BytesIO`` whose ``fileno()`` raises
# ``UnsupportedOperation``, so the DoS mitigation's ``select()`` guard is
# never exercised by transport tests routed through the mock. These steps
# wire a real ``os.pipe()`` into the server so the actual timeout code path
# (the one that protects production from partial-body stalls) runs end to end.
@when(
"I invoke read body with timeout {timeout:f} on an empty pipe expecting "
"{expected:d} bytes"
)
def step_read_body_pipe_no_data(
context: Context, timeout: float, expected: int
) -> None:
"""Exercise ``_read_body_with_timeout``'s select-based timeout path.
A pipe fd makes ``fileno()`` succeed so ``use_select`` is True and the
DoS-protection branch runs. No bytes are written before the call:
* ``timeout=0.0`` makes the deadline already past on entry, exercising
the ``if timeout <= 0:`` early-exit warning.
* A small positive timeout (e.g. ``0.05``) lets ``select.select()`` run
and return no ready descriptors, exercising the ``if not ready:``
warning.
Both paths log ``lsp.transport.read_timeout`` and return ``None``.
"""
read_fd, write_fd = os.pipe()
pipe_in = os.fdopen(read_fd, "rb", buffering=0)
server = LspServer(
input_stream=pipe_in,
output_stream=io.BytesIO(),
read_timeout=timeout,
)
try:
with _capture_structlogs() as captured:
result = server._read_body_with_timeout(expected)
context.lsp_read_body_result = result
context.lsp_captured_logs = list(captured)
finally:
os.close(write_fd)
pipe_in.close()
@then("the read body result should be None")
def step_read_body_result_none(context: Context) -> None:
"""Assert the timeout path returned ``None`` (so the caller can ``_SKIP``)."""
assert context.lsp_read_body_result is None, (
f"expected None, got {context.lsp_read_body_result!r}"
)
@then("a transport read_timeout warning should be captured")
def step_transport_read_timeout_captured(context: Context) -> None:
"""Assert at least one ``lsp.transport.read_timeout`` event was logged."""
matches = [
entry
for entry in context.lsp_captured_logs
if entry.get("event") == "lsp.transport.read_timeout"
]
assert matches, (
"no lsp.transport.read_timeout warning in captured logs: "
f"{context.lsp_captured_logs}"
)