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
16 changes: 16 additions & 0 deletions roboflow/cli/_output.py
Original file line number Diff line number Diff line change
Expand Up @@ -281,3 +281,19 @@ def suppress_sdk_output(args: Any = None) -> Iterator[None]:
"""
with contextlib.redirect_stdout(io.StringIO()):
yield


@contextlib.contextmanager
def sdk_output_to_stderr(args: Any) -> Iterator[None]:
"""Move SDK stdout output to stderr in ``--json`` mode.

Some SDK calls print progress or per-item status that is worth seeing,
for example the per-image results of a directory upload, so it is not
suppressed. With ``--json`` it goes to stderr, so that stdout carries
only the JSON document written by ``output()``.
"""
if getattr(args, "json", False):
with contextlib.redirect_stdout(sys.stderr):
yield
else:
yield
29 changes: 15 additions & 14 deletions roboflow/cli/handlers/image.py
Original file line number Diff line number Diff line change
Expand Up @@ -359,7 +359,7 @@ def _handle_upload_directory(args, api_key: str, path: str) -> None: # noqa: AN
import os

import roboflow
from roboflow.cli._output import output, output_error, suppress_sdk_output
from roboflow.cli._output import output, output_error, sdk_output_to_stderr, suppress_sdk_output

# Always suppress SDK "loading..." noise during workspace init
with suppress_sdk_output():
Expand All @@ -376,19 +376,20 @@ def _handle_upload_directory(args, api_key: str, path: str) -> None: # noqa: AN
wait = not getattr(args, "no_wait", False)

try:
result = workspace.upload_dataset(
dataset_path=path,
project_name=args.project,
num_workers=args.concurrency,
batch_name=getattr(args, "batch", None),
num_retries=retries,
is_prediction=getattr(args, "is_prediction", False),
use_zip_upload=getattr(args, "zip_upload", False),
annotation_overwrite=getattr(args, "annotation_overwrite", None),
split=getattr(args, "split", None),
tags=tags,
wait=wait,
)
with sdk_output_to_stderr(args):
result = workspace.upload_dataset(
dataset_path=path,
project_name=args.project,
num_workers=args.concurrency,
batch_name=getattr(args, "batch", None),
num_retries=retries,
is_prediction=getattr(args, "is_prediction", False),
use_zip_upload=getattr(args, "zip_upload", False),
annotation_overwrite=getattr(args, "annotation_overwrite", None),
split=getattr(args, "split", None),
tags=tags,
wait=wait,
)
except Exception as exc:
output_error(args, str(exc))
return
Expand Down
30 changes: 17 additions & 13 deletions roboflow/cli/handlers/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -311,11 +311,13 @@ def _get_model(args): # noqa: ANN001

def _upload_model(args): # noqa: ANN001
import roboflow
from roboflow.cli._output import output, output_error
from roboflow.cli._output import output, output_error, sdk_output_to_stderr, suppress_sdk_output

api_key = args.api_key or None
rf = roboflow.Roboflow(api_key=api_key)
workspace = rf.workspace(args.workspace)
# Always suppress SDK "loading..." noise during workspace init
with suppress_sdk_output():
rf = roboflow.Roboflow(api_key=api_key)
workspace = rf.workspace(args.workspace)

if args.version_number is not None:
# Deploy to a specific version
Expand All @@ -325,9 +327,10 @@ def _upload_model(args): # noqa: ANN001
return

try:
project = workspace.project(project_id)
version = project.version(args.version_number)
version.deploy(str(args.model_type), str(args.model_path), str(args.filename))
with sdk_output_to_stderr(args):
project = workspace.project(project_id)
version = project.version(args.version_number)
version.deploy(str(args.model_type), str(args.model_path), str(args.filename))
except Exception as exc:
output_error(args, str(exc))
return
Expand All @@ -338,13 +341,14 @@ def _upload_model(args): # noqa: ANN001
return

try:
workspace.deploy_model(
model_type=str(args.model_type),
model_path=str(args.model_path),
project_ids=args.project,
model_name=str(args.model_name) if args.model_name else "",
filename=str(args.filename),
)
with sdk_output_to_stderr(args):
workspace.deploy_model(
model_type=str(args.model_type),
model_path=str(args.model_path),
project_ids=args.project,
model_name=str(args.model_name) if args.model_name else "",
filename=str(args.filename),
)
except Exception as exc:
output_error(args, str(exc))
return
Expand Down
21 changes: 11 additions & 10 deletions roboflow/cli/handlers/search.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,18 +149,19 @@ def _describe_hit(hit: dict) -> str:


def _do_export(args: Any, workspace: Any) -> None:
from roboflow.cli._output import output, output_error
from roboflow.cli._output import output, output_error, sdk_output_to_stderr

try:
result_path = workspace.search_export(
query=args.query,
format=args.format,
location=args.location,
dataset=args.dataset,
annotation_group=getattr(args, "annotation_group", None),
name=args.name,
extract_zip=not args.no_extract,
)
with sdk_output_to_stderr(args):
result_path = workspace.search_export(
query=args.query,
format=args.format,
location=args.location,
dataset=args.dataset,
annotation_group=getattr(args, "annotation_group", None),
name=args.name,
extract_zip=not args.no_extract,
)
except Exception as exc:
output_error(args, str(exc))
return
Expand Down
5 changes: 3 additions & 2 deletions roboflow/cli/handlers/version.py
Original file line number Diff line number Diff line change
Expand Up @@ -230,7 +230,7 @@ def _parse_url(url: str) -> tuple:

def _download(args): # noqa: ANN001
import roboflow
from roboflow.cli._output import output, output_error, suppress_sdk_output
from roboflow.cli._output import output, output_error, sdk_output_to_stderr, suppress_sdk_output

w, p, v = _parse_url(args.url_or_id)

Expand All @@ -257,7 +257,8 @@ def _download(args): # noqa: ANN001
else:
version_obj = project.version(int(v))

version_obj.download(args.format, location=args.location, overwrite=True)
with sdk_output_to_stderr(args):
version_obj.download(args.format, location=args.location, overwrite=True)
except SystemExit:
raise
except Exception as exc:
Expand Down
134 changes: 134 additions & 0 deletions tests/cli/test_json_stdout.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,134 @@
"""stdout stays valid JSON in --json mode when SDK calls print progress."""

import contextlib
import io
import json
import os
import tempfile
import types
import unittest
from unittest.mock import MagicMock, patch


def _sdk_print(*_args: object, **_kwargs: object) -> None:
print("progress line printed by the SDK")


def _run(handler, *args: object) -> tuple: # noqa: ANN001
stdout, stderr = io.StringIO(), io.StringIO()
with contextlib.redirect_stdout(stdout), contextlib.redirect_stderr(stderr):
handler(*args)
return stdout.getvalue(), stderr.getvalue()


class TestJsonStdout(unittest.TestCase):
def assert_json_stdout(self, stdout: str, stderr: str) -> dict:
self.assertNotIn("progress line printed by the SDK", stdout)
self.assertIn("progress line printed by the SDK", stderr)
return json.loads(stdout)

def test_search_export(self) -> None:
from roboflow.cli.handlers.search import _do_export

def search_export(**_kwargs: object) -> str:
_sdk_print()
return "export-dir"

workspace = MagicMock()
workspace.search_export.side_effect = search_export
args = types.SimpleNamespace(
json=True,
query="tag:a",
format="coco",
location="export-dir",
dataset=None,
annotation_group=None,
name=None,
no_extract=False,
)

stdout, stderr = _run(_do_export, args, workspace)

self.assertEqual(self.assert_json_stdout(stdout, stderr)["status"], "completed")

@patch("roboflow.Roboflow")
def test_image_upload_directory(self, mock_rf_cls: MagicMock) -> None:
from roboflow.cli.handlers.image import _handle_upload_directory

mock_rf_cls.return_value.workspace.return_value.upload_dataset.side_effect = _sdk_print
with tempfile.TemporaryDirectory() as tmpdir:
with open(os.path.join(tmpdir, "a.jpg"), "w") as f:
f.write("x")
args = types.SimpleNamespace(
json=True,
workspace="ws",
project="proj",
concurrency=1,
retries=0,
tag=None,
batch=None,
split=None,
is_prediction=False,
zip_upload=False,
annotation_overwrite=None,
no_wait=False,
)

stdout, stderr = _run(_handle_upload_directory, args, "key", tmpdir)

self.assertEqual(self.assert_json_stdout(stdout, stderr)["count"], 1)

@patch("roboflow.Roboflow")
def test_model_upload(self, mock_rf_cls: MagicMock) -> None:
from roboflow.cli.handlers.model import _upload_model

def workspace(*_args: object) -> MagicMock:
print("loading Roboflow workspace...")
mock_workspace = MagicMock()
mock_workspace.project.return_value.version.return_value.deploy.side_effect = _sdk_print
return mock_workspace

mock_rf_cls.return_value.workspace.side_effect = workspace
args = types.SimpleNamespace(
json=True,
api_key="key",
workspace="ws",
project=["proj"],
version_number=1,
model_type="yolov8",
model_path="/path/to/model",
filename="weights/best.pt",
model_name=None,
)

stdout, stderr = _run(_upload_model, args)

self.assertEqual(self.assert_json_stdout(stdout, stderr)["status"], "uploaded")
self.assertNotIn("loading Roboflow workspace", stdout + stderr)

@patch("roboflow.Roboflow")
def test_version_download(self, mock_rf_cls: MagicMock) -> None:
from roboflow.cli.handlers.version import _download

project = mock_rf_cls.return_value.workspace.return_value.project.return_value
project.version.return_value.download.side_effect = _sdk_print
args = types.SimpleNamespace(json=True, url_or_id="ws/proj/1", format="coco", location="dataset-dir")

stdout, stderr = _run(_download, args)

self.assertEqual(self.assert_json_stdout(stdout, stderr)["version"], 1)

def test_text_mode_keeps_sdk_output_on_stdout(self) -> None:
from roboflow.cli._output import sdk_output_to_stderr

stdout, stderr = io.StringIO(), io.StringIO()
with contextlib.redirect_stdout(stdout), contextlib.redirect_stderr(stderr):
with sdk_output_to_stderr(types.SimpleNamespace(json=False)):
_sdk_print()

self.assertIn("progress line printed by the SDK", stdout.getvalue())
self.assertEqual(stderr.getvalue(), "")


if __name__ == "__main__":
unittest.main()
Loading