Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 21 additions & 2 deletions pyk/src/pyk/cterm/symbolic.py
Original file line number Diff line number Diff line change
Expand Up @@ -314,6 +314,16 @@ def add_module(self, module: KFlatModule, name_as_id: bool = False) -> str:
_kore_module = kflatmodule_to_kore(self._definition, module)
return self._kore_client.add_module(_kore_module, name_as_id=name_as_id)

def set_client_label(self, label: str) -> None:
"""Stamp `label` on every subsequent JSON-RPC request issued through this client.

Called automatically by APRProver/ImpliesProver at the start of each
proof so that booster's per-line `{request: ...}` context self-identifies
which claim drove the work. Safe to call between proofs to re-tag the
same client for a different claim.
"""
self._kore_client.set_client_label(label)

def _smt_solver_error(self, err: SmtSolverError) -> CTermSMTError:
kast = self.kore_to_kast(err.pattern)
pretty_pattern = PrettyPrinter(self._definition).print(kast)
Expand All @@ -326,6 +336,7 @@ def cterm_symbolic(
definition_dir: Path,
*,
id: str | None = None,
client_label: str | None = None,
port: int | None = None,
kore_rpc_command: str | Iterable[str] | None = None,
llvm_definition_dir: Path | None = None,
Expand Down Expand Up @@ -365,7 +376,13 @@ def cterm_symbolic(
simplify_each=simplify_each,
no_post_exec_simplify=no_post_exec_simplify,
) as server:
with KoreClient('localhost', server.port, bug_report=bug_report, bug_report_id=id) as client:
with KoreClient(
'localhost',
server.port,
bug_report=bug_report,
bug_report_id=id,
client_label=client_label,
) as client:
yield CTermSymbolic(
client,
definition,
Expand All @@ -376,7 +393,9 @@ def cterm_symbolic(
else:
if port is None:
raise ValueError('Missing port with start_server=False')
with KoreClient('localhost', port, bug_report=bug_report, bug_report_id=id) as client:
with KoreClient(
'localhost', port, bug_report=bug_report, bug_report_id=id, client_label=client_label
) as client:
yield CTermSymbolic(
client, definition, log_succ_rewrites=log_succ_rewrites, log_fail_rewrites=log_fail_rewrites
)
53 changes: 51 additions & 2 deletions pyk/src/pyk/kore/rpc.py
Original file line number Diff line number Diff line change
Expand Up @@ -193,6 +193,7 @@ def __init__(
timeout: int | None = None,
bug_report: BugReport | None = None,
bug_report_id: str | None = None,
client_label: str | None = None,
):
client_cache = {}
self._clients = {}
Expand All @@ -202,6 +203,7 @@ def __init__(
timeout=timeout,
bug_report=bug_report,
bug_report_id=bug_report_id,
client_label=client_label,
transport=default_transport,
)
client_cache[(default_host, default_port)] = self._default_client
Expand All @@ -212,7 +214,13 @@ def __init__(
else:
new_id = None if bug_report_id is None else bug_report_id + '_' + str(transport)
new_client = JsonRpcClient(
host, port, timeout=timeout, bug_report=bug_report, bug_report_id=new_id, transport=transport
host,
port,
timeout=timeout,
bug_report=bug_report,
bug_report_id=new_id,
client_label=client_label,
transport=transport,
)
self._update_clients(method, new_client)
client_cache[(host, port)] = new_client
Expand All @@ -237,6 +245,23 @@ def close(self) -> None:
for client in clients:
client.close()

def _unique_clients(self) -> list[JsonRpcClient]:
# Dispatch lists may share clients across methods (one JsonRpcClient per
# distinct (host, port)); collect each client only once by identity.
result = [self._default_client]
seen = {id(self._default_client)}
for clients in self._clients.values():
for client in clients:
if id(client) not in seen:
seen.add(id(client))
result.append(client)
return result

def set_client_label(self, label: str) -> None:
"""Push a new client_label onto every distinct underlying JsonRpcClient."""
for client in self._unique_clients():
client.set_client_label(label)

def request(self, method: str, **params: Any) -> dict[str, Any]:
if method in self._clients:
for client in self._clients[method]:
Expand All @@ -253,6 +278,7 @@ class JsonRpcClient(ContextManager['JsonRpcClient']):

_transport: Transport
_req_id: int
_client_label: str

_bug_report: BugReport | None
_bug_report_id: str | None
Expand All @@ -265,12 +291,17 @@ def __init__(
timeout: int | None = None,
bug_report: BugReport | None = None,
bug_report_id: str | None = None,
client_label: str | None = None,
transport: TransportType = TransportType.SINGLE_SOCKET,
):
self._transport = self._create_transport(transport, host=host, port=port, timeout=timeout)
self._req_id = 1
self._bug_report_id = bug_report_id
self._bug_report = bug_report
# Stamped on every outgoing request-id as `{client_label}-NNN` so that
# booster's `{request: ...}` context lines self-identify the caller.
# Defaults to str(id(self)) so non-adopters keep the prior shape byte-for-byte.
self._client_label = client_label if client_label is not None else str(id(self))

@staticmethod
def _create_transport(transport: TransportType, *, host: str, port: int, timeout: int | None) -> Transport:
Expand All @@ -292,7 +323,7 @@ def close(self) -> None:
self._transport.close()

def request(self, method: str, **params: Any) -> dict[str, Any]:
req_id = f'{id(self)}-{self._req_id:03}'
req_id = f'{self._client_label}-{self._req_id:03}'
self._req_id += 1

payload = {
Expand Down Expand Up @@ -334,6 +365,14 @@ def _check(response: Mapping[str, Any]) -> None:
assert response['error']['code'] not in {-32700, -32600}, 'Malformed JSON-RPC request'
raise JsonRpcError(**response['error'])

def set_client_label(self, label: str) -> None:
"""Set the prefix stamped on every subsequent request id.

Persists until the next `set_client_label` (or close). Used by
APRProver/ImpliesProver to re-tag a shared client per claim.
"""
self._client_label = label


class KoreClientError(Exception, ABC):
def __init__(self, message: str):
Expand Down Expand Up @@ -895,6 +934,7 @@ def __init__(
timeout: int | None = None,
bug_report: BugReport | None = None,
bug_report_id: str | None = None,
client_label: str | None = None,
transport: TransportType = TransportType.SINGLE_SOCKET,
dispatch: dict[str, list[tuple[str, int, TransportType]]] | None = None,
):
Expand All @@ -908,9 +948,18 @@ def __init__(
timeout=timeout,
bug_report=bug_report,
bug_report_id=bug_report_id,
client_label=client_label,
dispatch=dispatch,
)

def set_client_label(self, label: str) -> None:
"""Set the label stamped on every subsequent JSON-RPC request id.

Persists until the next `set_client_label` (or close). Booster's
per-line `{request: <id>}` context then self-identifies the caller.
"""
self._client.set_client_label(label)

def __enter__(self) -> KoreClient:
return self

Expand Down
5 changes: 4 additions & 1 deletion pyk/src/pyk/proof/implies.py
Original file line number Diff line number Diff line change
Expand Up @@ -466,7 +466,10 @@ def step_proof(self, step: ImpliesProofStep) -> list[ImpliesProofResult]:
]

def init_proof(self, proof: ImpliesProof) -> None:
pass
# Stamp proof.id on every subsequent kore-RPC request so booster's
# per-line `{request: ...}` context self-identifies the claim driving
# the work — no PID join through pyk.log needed downstream.
self.kcfg_explore.cterm_symbolic.set_client_label(proof.id)

def failure_info(self, proof: ImpliesProof) -> FailureInfo:
# TODO add implementation
Expand Down
4 changes: 4 additions & 0 deletions pyk/src/pyk/proof/reachability.py
Original file line number Diff line number Diff line change
Expand Up @@ -756,6 +756,10 @@ def close(self) -> None:
self.kcfg_explore.cterm_symbolic._kore_client.close()

def init_proof(self, proof: APRProof) -> None:
# Stamp proof.id on every subsequent kore-RPC request so booster's
# per-line `{request: ...}` context self-identifies the claim driving
# the work — no PID join through pyk.log needed downstream.
self.kcfg_explore.cterm_symbolic.set_client_label(proof.id)
main_module_name = self.main_module_name
if self.extra_module:
main_module_name = self.kcfg_explore.cterm_symbolic.add_module(self.extra_module, name_as_id=True)
Expand Down
8 changes: 7 additions & 1 deletion pyk/src/tests/unit/kore/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,7 +90,13 @@ def rpc_client(mock: Mock) -> MockClient:
def kore_client(mock: Mock, mock_class: Mock) -> Iterator[KoreClient]: # noqa: N803
client = KoreClient('localhost', 3000)
mock_class.assert_called_with(
'localhost', 3000, timeout=None, bug_report=None, bug_report_id=None, transport=TransportType.SINGLE_SOCKET
'localhost',
3000,
timeout=None,
bug_report=None,
bug_report_id=None,
client_label=None,
transport=TransportType.SINGLE_SOCKET,
)
assert client._client._default_client == mock
yield client
Expand Down
118 changes: 118 additions & 0 deletions pyk/src/tests/unit/kore/test_client_label.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
"""Unit tests for the caller-supplied `client_label` on JsonRpcClient / KoreClient.

The label is stamped on every outgoing request-id (`{label}-NNN`) so booster's
per-line `{request: ...}` context self-identifies the caller (typically the
claim being discharged). `set_client_label(label)` mutates the prefix for all
subsequent requests; in normal pyk use, APRProver/ImpliesProver invoke this
automatically per proof so consumers never touch the API directly.
"""

from __future__ import annotations

import json
from typing import TYPE_CHECKING
from unittest.mock import patch

import pytest

from pyk.kore.rpc import JsonRpcClient, KoreClient, SingleSocketTransport

if TYPE_CHECKING:
from collections.abc import Iterator
from unittest.mock import Mock


@pytest.fixture
def mock_class() -> Iterator[Mock]:
patcher = patch('pyk.kore.rpc.SingleSocketTransport', spec=True)
yield patcher.start()
patcher.stop()


@pytest.fixture
def mock(mock_class: Mock) -> Mock:
mock = mock_class.return_value
assert isinstance(mock, SingleSocketTransport)
return mock # type: ignore


def _wire_capture(mock: Mock) -> list[dict]:
"""Capture every outgoing payload and echo a success response."""
captured: list[dict] = []

def respond(req: str, req_id: str, method_name: str) -> str:
payload = json.loads(req)
captured.append(payload)
return json.dumps({'jsonrpc': '2.0', 'id': payload['id'], 'result': {}})

mock.request.side_effect = respond
return captured


def test_default_label_uses_object_id(mock: Mock) -> None:
"""Without a client_label, the request-id prefix is str(id(self)) — byte-stable with the legacy path."""
captured = _wire_capture(mock)
client = JsonRpcClient('localhost', 3000)
expected_prefix = str(id(client))

client.request('execute')
client.request('simplify')

assert [p['id'] for p in captured] == [f'{expected_prefix}-001', f'{expected_prefix}-002']


def test_construction_label_is_used_as_prefix(mock: Mock) -> None:
captured = _wire_capture(mock)
client = JsonRpcClient('localhost', 3000, client_label='LEMMAS-SPEC.range-31')

client.request('execute')
client.request('simplify')

assert [p['id'] for p in captured] == ['LEMMAS-SPEC.range-31-001', 'LEMMAS-SPEC.range-31-002']


def test_set_client_label_swaps_prefix_for_subsequent_requests(mock: Mock) -> None:
captured = _wire_capture(mock)
client = JsonRpcClient('localhost', 3000, client_label='claim-A')

client.request('execute')
client.set_client_label('claim-B')
client.request('simplify')
client.request('implies')

assert [p['id'] for p in captured] == ['claim-A-001', 'claim-B-002', 'claim-B-003']


def test_set_client_label_persists_no_restoration(mock: Mock) -> None:
"""The setter is permanent — there is no enclosing scope or restoration semantics."""
captured = _wire_capture(mock)
client = JsonRpcClient('localhost', 3000, client_label='outer')

client.set_client_label('inner')
client.request('execute')
client.request('simplify')

assert [p['id'] for p in captured] == ['inner-001', 'inner-002']


def test_kore_client_forwards_client_label(mock: Mock) -> None:
"""KoreClient(client_label=...) plumbs through to the underlying JsonRpcClient."""
captured = _wire_capture(mock)
kore_client = KoreClient('localhost', 3000, client_label='claim-X')

assert kore_client._client._default_client._client_label == 'claim-X'

kore_client._client._default_client.request('execute')
assert captured[-1]['id'] == 'claim-X-001'


def test_kore_client_set_client_label_forwards_to_underlying_clients(mock: Mock) -> None:
captured = _wire_capture(mock)
kore_client = KoreClient('localhost', 3000, client_label='construction-default')

kore_client.set_client_label('claim-A')
kore_client._client._default_client.request('execute')
kore_client.set_client_label('claim-B')
kore_client._client._default_client.request('simplify')

assert [p['id'] for p in captured] == ['claim-A-001', 'claim-B-002']
Loading
Loading