Skip to content
Merged
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: 18 additions & 5 deletions python/cuopt/cuopt/grpc/client/grpc_client.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -970,13 +970,26 @@ cdef class RoutingClient:

cdef unique_ptr[grpc_python_client_t] _client

def __cinit__(self, str target="localhost:50051"):
host, _, port = target.rpartition(":")
if not host:
host, port = target, "50051"
def __cinit__(self, str host, int port, *, tls=None):
Comment thread
coderabbitai[bot] marked this conversation as resolved.
"""
Connect to ``cuopt_grpc_server`` at ``host:port``.

``tls`` controls transport security, same as :class:`Client`:

* ``None`` (default) — read ``CUOPT_TLS_*`` from the environment.
* ``False`` — plain TCP; ignore ``CUOPT_TLS_*``.
* :class:`TlsConfig` — explicit TLS/mTLS; omit ``root_certs`` to use the
system/default CA trust store.
"""
if tls is not None and tls is not False and not isinstance(tls, TlsConfig):
raise TypeError("tls must be None, False, or TlsConfig")

cdef grpc_python_client_connect_options_t options
cdef string host_cpp = host.encode("utf-8")
cdef string err
self._client.reset(new grpc_python_client_t(host_cpp, int(port)))

options = _connect_options_from_tls(tls)
self._client.reset(new grpc_python_client_t(host_cpp, port, options))
if not self._client.get().connect(err):
raise RoutingSolveError(
"failed to connect: " + err.decode("utf-8")
Expand Down
2 changes: 1 addition & 1 deletion python/cuopt/cuopt/grpc/routing/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
dm = routing.DataModel(n_locations, n_fleet)
dm.add_cost_matrix(cost)
...
client = RoutingClient("gpu-host:50051")
client = RoutingClient("gpu-host", 50051)
solution = client.solve(dm)
"""

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,11 @@
)


def _client():
host, _, port = (_SERVER or "").rpartition(":")
return grpc_routing.RoutingClient(host, int(port))


def _small_vrp():
dm = routing.DataModel(5, 2)
cost = np.array(
Expand All @@ -46,7 +51,7 @@ def test_remote_solve_matches_local():
settings.set_time_limit(2)
local = routing.Solve(_small_vrp(), settings)

client = grpc_routing.RoutingClient(_SERVER)
client = _client()
remote = client.solve(_small_vrp(), {"time_limit": 2.0})

assert remote["status"] == 0, remote["status_message"]
Expand All @@ -57,7 +62,7 @@ def test_remote_solve_matches_local():


def test_submit_wait_result_lifecycle():
client = grpc_routing.RoutingClient(_SERVER)
client = _client()
job_id = client.submit(_small_vrp(), {"time_limit": 1.0})
assert job_id
client.wait(job_id, timeout=30)
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,84 @@
# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""TLS/mTLS coverage for the compiled VRP gRPC client (cuopt.grpc.routing).

Unlike test_routing_grpc_client.py, these tests do not need
CUOPT_GRPC_SERVER -- the tls_server_info/mtls_server_info fixtures start
their own cuopt_grpc_server subprocess (skipped if the test TLS certs are
not found).
"""

import os

import numpy as np
import pytest

from cuopt import routing
from cuopt.grpc.linear_programming import TlsConfig

grpc_routing = pytest.importorskip("cuopt.grpc.routing")


def _small_vrp():
dm = routing.DataModel(5, 2)
cost = np.array(
[
[0, 1, 2, 2, 1],
[1, 0, 1, 2, 2],
[2, 1, 0, 1, 2],
[2, 2, 1, 0, 1],
[1, 2, 2, 1, 0],
],
dtype=np.float32,
)
dm.add_cost_matrix(cost)
return dm


def test_rejects_invalid_tls_argument():
with pytest.raises(TypeError):
grpc_routing.RoutingClient("localhost", 1, tls="bogus")


@pytest.mark.xdist_group(name="grpc_server")
@pytest.mark.filterwarnings("ignore::DeprecationWarning")
class TestRoutingClientTls:
def test_submit_with_explicit_tls_config(self, tls_server_info):
cert_dir = tls_server_info["cert_dir"]
client = grpc_routing.RoutingClient(
"localhost",
tls_server_info["port"],
tls=TlsConfig(os.path.join(cert_dir, "ca.crt")),
)
solution = client.solve(_small_vrp(), {"time_limit": 2.0})
assert solution["status"] == 0, solution["status_message"]

def test_tls_server_rejects_plain_client(self, tls_server_info):
with pytest.raises(grpc_routing.RoutingSolveError):
grpc_routing.RoutingClient(
"localhost", tls_server_info["port"], tls=False
)

def test_submit_with_explicit_mtls_config(self, mtls_server_info):
cert_dir = mtls_server_info["cert_dir"]
client = grpc_routing.RoutingClient(
"localhost",
mtls_server_info["port"],
tls=TlsConfig(
os.path.join(cert_dir, "ca.crt"),
client_cert=os.path.join(cert_dir, "client.crt"),
client_key=os.path.join(cert_dir, "client.key"),
),
)
solution = client.solve(_small_vrp(), {"time_limit": 2.0})
assert solution["status"] == 0, solution["status_message"]

def test_mtls_server_rejects_missing_client_cert(self, mtls_server_info):
cert_dir = mtls_server_info["cert_dir"]
with pytest.raises(grpc_routing.RoutingSolveError):
grpc_routing.RoutingClient(
"localhost",
mtls_server_info["port"],
tls=TlsConfig(os.path.join(cert_dir, "ca.crt")),
)
Loading