Skip to content
Open
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
17 changes: 13 additions & 4 deletions cuda_bindings/cuda/bindings/_example_helpers/helper_string.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,11 +5,20 @@


def check_cmd_line_flag(string_ref):
return any(string_ref == i and k < len(sys.argv) - 1 for i, k in enumerate(sys.argv))
"""Return whether ``string_ref`` was passed on the command line.

``sys.argv[0]`` is the program name and is never considered a flag.
"""
return string_ref in sys.argv[1:]


def get_cmd_line_argument_int(string_ref):
for i, k in enumerate(sys.argv):
if string_ref == i and k < len(sys.argv) - 1:
return sys.argv[k + 1]
"""Return the integer that follows ``string_ref`` on the command line.

Returns 0 if ``string_ref`` was not passed, or if nothing follows it.
"""
args = sys.argv[1:]
for idx, arg in enumerate(args):
if arg == string_ref and idx + 1 < len(args):
return int(args[idx + 1])
return 0
64 changes: 64 additions & 0 deletions cuda_bindings/tests/test_example_helpers.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

import sys

import pytest

from cuda.bindings._example_helpers import check_cmd_line_flag, get_cmd_line_argument_int


@pytest.fixture
def argv(monkeypatch):
"""Replace sys.argv, keeping a realistic program name in argv[0]."""

def _argv(*args, prog="example.py"):
monkeypatch.setattr(sys, "argv", [prog, *args])

return _argv


@pytest.mark.agent_authored(model="claude-opus-5")
@pytest.mark.parametrize(
("args", "flag", "expected"),
[
((), "device=", False),
(("device=", "0"), "device=", True),
(("help",), "help", True), # a boolean flag has nothing after it
(("wA=", "128", "hA=", "256"), "hA=", True),
(("wA=", "128"), "hA=", False),
],
)
def test_check_cmd_line_flag(argv, args, flag, expected):
argv(*args)
assert check_cmd_line_flag(flag) is expected


@pytest.mark.agent_authored(model="claude-opus-5")
def test_check_cmd_line_flag_ignores_the_program_name(argv):
argv(prog="help")
assert check_cmd_line_flag("help") is False


@pytest.mark.agent_authored(model="claude-opus-5")
@pytest.mark.parametrize(
("args", "expected"),
[
((), 0),
(("device=", "3"), 3),
(("wA=", "128", "device=", "2"), 2),
(("device=",), 0), # nothing follows the flag
(("nomatch", "7"), 0),
],
)
def test_get_cmd_line_argument_int(argv, args, expected):
argv(*args)
value = get_cmd_line_argument_int("device=")
assert value == expected
assert isinstance(value, int)


@pytest.mark.agent_authored(model="claude-opus-5")
def test_get_cmd_line_argument_int_ignores_the_program_name(argv):
argv("3", prog="device=")
assert get_cmd_line_argument_int("device=") == 0
Loading