cvat-cli 2.69.0__tar.gz → 2.71.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (29) hide show
  1. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/PKG-INFO +2 -2
  2. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/pyproject.toml +1 -1
  3. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/__main__.py +15 -3
  4. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/agent.py +4 -0
  5. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/agent_driver.py +24 -3
  6. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/agent_driver_detection.py +44 -4
  7. cvat_cli-2.71.0/src/cvat_cli/_internal/agent_driver_interaction.py +218 -0
  8. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/agent_driver_tracking.py +5 -2
  9. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/command_base.py +4 -1
  10. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/commands_all.py +2 -0
  11. cvat_cli-2.71.0/src/cvat_cli/_internal/commands_config.py +45 -0
  12. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/commands_functions.py +5 -47
  13. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/common.py +15 -99
  14. cvat_cli-2.71.0/src/cvat_cli/version.py +1 -0
  15. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli.egg-info/PKG-INFO +2 -2
  16. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli.egg-info/SOURCES.txt +2 -0
  17. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli.egg-info/requires.txt +1 -1
  18. cvat_cli-2.69.0/src/cvat_cli/version.py +0 -1
  19. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/README.md +0 -0
  20. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/setup.cfg +0 -0
  21. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/__init__.py +0 -0
  22. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/__init__.py +0 -0
  23. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/commands_projects.py +0 -0
  24. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/commands_tasks.py +0 -0
  25. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/parsers.py +0 -0
  26. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/utils.py +0 -0
  27. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli.egg-info/dependency_links.txt +0 -0
  28. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli.egg-info/entry_points.txt +0 -0
  29. {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli.egg-info/top_level.txt +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: cvat-cli
3
- Version: 2.69.0
3
+ Version: 2.71.0
4
4
  Summary: Command-line client for CVAT
5
5
  Author-email: "CVAT.ai Corporation" <support@cvat.ai>
6
6
  License-Expression: MIT
@@ -9,7 +9,7 @@ Classifier: Programming Language :: Python :: 3
9
9
  Classifier: Operating System :: OS Independent
10
10
  Requires-Python: >=3.10
11
11
  Description-Content-Type: text/markdown
12
- Requires-Dist: cvat-sdk==2.69.0
12
+ Requires-Dist: cvat-sdk==2.71.0
13
13
  Requires-Dist: attrs>=24.2.0
14
14
  Requires-Dist: Pillow>=10.3.0
15
15
 
@@ -16,7 +16,7 @@ classifiers = [
16
16
  ]
17
17
  requires-python = ">=3.10"
18
18
  dependencies = [
19
- "cvat-sdk==2.69.0",
19
+ "cvat-sdk==2.71.0",
20
20
 
21
21
  "attrs>=24.2.0",
22
22
  "Pillow>=10.3.0",
@@ -9,6 +9,7 @@ import sys
9
9
 
10
10
  import urllib3.exceptions
11
11
  from cvat_sdk import exceptions
12
+ from cvat_sdk.core.exceptions import AuthStoreError
12
13
 
13
14
  from ._internal.commands_all import COMMANDS
14
15
  from ._internal.common import (
@@ -31,10 +32,21 @@ def main(args: list[str] = None):
31
32
 
32
33
  configure_logger(logger, parsed_args)
33
34
 
35
+ executor = popattr(parsed_args, "_executor")
36
+ needs_client = popattr(parsed_args, "_needs_client")
37
+
34
38
  try:
35
- with build_client(parsed_args, logger=logger) as client:
36
- popattr(parsed_args, "_executor")(client, **vars(parsed_args))
37
- except (exceptions.ApiException, urllib3.exceptions.HTTPError, CriticalError) as e:
39
+ if needs_client:
40
+ with build_client(parsed_args, logger=logger) as client:
41
+ executor(client, **vars(parsed_args))
42
+ else:
43
+ executor(parsed_args)
44
+ except (
45
+ exceptions.ApiException,
46
+ urllib3.exceptions.HTTPError,
47
+ CriticalError,
48
+ AuthStoreError,
49
+ ) as e:
38
50
  logger.critical(e)
39
51
  return 1
40
52
 
@@ -36,6 +36,7 @@ from .agent_driver import (
36
36
  worker_current_function,
37
37
  )
38
38
  from .agent_driver_detection import AgentDetectionFunctionDriver
39
+ from .agent_driver_interaction import AgentInteractionFunctionDriver
39
40
  from .agent_driver_tracking import AgentTrackingFunctionDriver, TrackingStateIdGenerator
40
41
  from .common import CriticalError, FunctionLoader
41
42
 
@@ -227,6 +228,9 @@ def get_function_driver_class(function_spec: object) -> type[AgentFunctionDriver
227
228
  if isinstance(function_spec, cvataa.DetectionFunctionSpec):
228
229
  return AgentDetectionFunctionDriver
229
230
 
231
+ if isinstance(function_spec, cvataa.InteractionFunctionSpec):
232
+ return AgentInteractionFunctionDriver
233
+
230
234
  if isinstance(function_spec, cvataa.TrackingFunctionSpec):
231
235
  return AgentTrackingFunctionDriver
232
236
 
@@ -5,7 +5,7 @@
5
5
  from __future__ import annotations
6
6
 
7
7
  from collections.abc import Callable
8
- from typing import TYPE_CHECKING, Any, ClassVar, Protocol
8
+ from typing import TYPE_CHECKING, Any, ClassVar, Generic, Protocol, TypeVar
9
9
 
10
10
  import cvat_sdk.auto_annotation as cvataa
11
11
  import cvat_sdk.datasets as cvatds
@@ -47,10 +47,13 @@ class CheckInCallback(Protocol):
47
47
  """
48
48
 
49
49
 
50
- class AgentFunctionDriver:
50
+ SpecT = TypeVar("SpecT")
51
+
52
+
53
+ class AgentFunctionDriver(Generic[SpecT]):
51
54
  FUNCTION_KIND: ClassVar[str]
52
55
 
53
- def __init__(self, client: Client, executor: RecoverableExecutor, function_spec: object):
56
+ def __init__(self, client: Client, executor: RecoverableExecutor, function_spec: SpecT) -> None:
54
57
  self._client = client
55
58
  self._executor = executor
56
59
  self._function_spec = function_spec
@@ -59,6 +62,10 @@ class AgentFunctionDriver:
59
62
  def init_worker(cls, state_id_generator: TrackingStateIdGenerator) -> None:
60
63
  pass
61
64
 
65
+ @classmethod
66
+ def get_remote_function_fields(cls, spec: SpecT) -> dict[str, Any]:
67
+ raise NotImplementedError
68
+
62
69
  def validate_function_compatibility(self, remote_function: dict) -> None:
63
70
  raise NotImplementedError
64
71
 
@@ -93,3 +100,17 @@ class AgentFunctionDriver:
93
100
  raise BadArError(f"Frame with index {frame_index} does not exist in the task")
94
101
 
95
102
  return sample, ds.labels
103
+
104
+ def _load_image_for_ar(self, sample, ar_params):
105
+ image = sample.media.load_image()
106
+
107
+ if roi := ar_params.get("roi"):
108
+ xtl, ytl, xbr, ybr = roi
109
+ width, height = image.size
110
+
111
+ if xbr > width or ybr > height:
112
+ raise BadArError("Invalid ROI")
113
+
114
+ return image.crop((xtl, ytl, xbr, ybr))
115
+
116
+ return image
@@ -30,9 +30,45 @@ def _worker_job_detect(
30
30
  return current_function.detect(context, image)
31
31
 
32
32
 
33
- class AgentDetectionFunctionDriver(AgentFunctionDriver):
33
+ class AgentDetectionFunctionDriver(AgentFunctionDriver[cvataa.DetectionFunctionSpec]):
34
34
  FUNCTION_KIND = "detector"
35
- _function_spec: cvataa.DetectionFunctionSpec
35
+
36
+ @staticmethod
37
+ def _dump_sublabel_spec(
38
+ sl_spec: models.SublabelRequest | models.PatchedLabelRequest,
39
+ ) -> dict:
40
+ result = {
41
+ "name": sl_spec.name,
42
+ "attributes": [
43
+ {
44
+ "name": attribute_spec.name,
45
+ "input_type": attribute_spec.input_type,
46
+ "values": attribute_spec.values,
47
+ }
48
+ for attribute_spec in getattr(sl_spec, "attributes", [])
49
+ ],
50
+ }
51
+
52
+ if getattr(sl_spec, "type", "any") != "any":
53
+ # Add the type conditionally, to stay compatible with older
54
+ # CVAT versions when the function doesn't define label types.
55
+ result["type"] = sl_spec.type
56
+
57
+ return result
58
+
59
+ @classmethod
60
+ def get_remote_function_fields(cls, spec: cvataa.DetectionFunctionSpec) -> dict[str, Any]:
61
+ labels_v2 = []
62
+
63
+ for label_spec in spec.labels:
64
+ labels_v2.append(cls._dump_sublabel_spec(label_spec))
65
+
66
+ if sublabels := getattr(label_spec, "sublabels", None):
67
+ labels_v2[-1]["sublabels"] = [
68
+ cls._dump_sublabel_spec(sublabel) for sublabel in sublabels
69
+ ]
70
+
71
+ return {"labels_v2": labels_v2}
36
72
 
37
73
  def _validate_sublabel_compatibility(
38
74
  self, remote_sl: dict, sl: models.Sublabel | None, sl_desc: str
@@ -134,7 +170,9 @@ class AgentDetectionFunctionDriver(AgentFunctionDriver):
134
170
  for sample_index, sample in enumerate(samples):
135
171
  context = self._create_detection_function_context(ar_params, sample.frame_name)
136
172
  annotations = self._executor.result(
137
- self._executor.submit(_worker_job_detect, context, sample.media.load_image())
173
+ self._executor.submit(
174
+ _worker_job_detect, context, self._load_image_for_ar(sample, ar_params)
175
+ )
138
176
  )
139
177
 
140
178
  tags, shapes = mapper.validate_and_remap(annotations, sample.frame_index)
@@ -155,7 +193,9 @@ class AgentDetectionFunctionDriver(AgentFunctionDriver):
155
193
  context = self._create_detection_function_context(ar_params, sample.frame_name)
156
194
 
157
195
  annotations = self._executor.result(
158
- self._executor.submit(_worker_job_detect, context, sample.media.load_image())
196
+ self._executor.submit(
197
+ _worker_job_detect, context, self._load_image_for_ar(sample, ar_params)
198
+ )
159
199
  )
160
200
 
161
201
  tags, shapes = mapper.validate_and_remap(annotations, sample.frame_index)
@@ -0,0 +1,218 @@
1
+ # Copyright (C) CVAT.ai Corporation
2
+ #
3
+ # SPDX-License-Identifier: MIT
4
+
5
+ from collections import OrderedDict
6
+ from datetime import datetime, timedelta, timezone
7
+ from typing import Any, cast
8
+
9
+ import attrs
10
+ import cvat_sdk.auto_annotation as cvataa
11
+ import PIL.Image
12
+
13
+ from .agent_driver import (
14
+ AgentFunctionDriver,
15
+ BadArError,
16
+ IncompatibleFunctionError,
17
+ worker_current_function,
18
+ )
19
+
20
+ _MAX_AGE_OF_PP_IMAGE = timedelta(minutes=10)
21
+
22
+
23
+ class _InteractionFunctionContextImpl(cvataa.InteractionFunctionContext):
24
+ pass
25
+
26
+
27
+ @attrs.frozen(kw_only=True)
28
+ class _InteractionPromptsImpl(cvataa.InteractionPrompts):
29
+ pos_points: tuple[tuple[float, float], ...]
30
+ neg_points: tuple[tuple[float, float], ...]
31
+ bounding_box: tuple[tuple[float, float], tuple[float, float]] | None
32
+
33
+
34
+ @attrs.define(kw_only=True)
35
+ class _PpImageCacheEntry:
36
+ pp_image: object
37
+ last_accessed_at: datetime = attrs.field(factory=lambda: datetime.now(tz=timezone.utc))
38
+
39
+
40
+ class _PpImageCache:
41
+ _MISSING = object()
42
+
43
+ def __init__(self):
44
+ self._contents = OrderedDict[object, _PpImageCacheEntry]()
45
+
46
+ def get(self, key: object) -> object:
47
+ if cache_entry := self._contents.get(key):
48
+ cache_entry.last_accessed_at = datetime.now(tz=timezone.utc)
49
+ self._contents.move_to_end(key)
50
+ return cache_entry.pp_image
51
+
52
+ return self._MISSING
53
+
54
+ def set(self, key: object, pp_image: object) -> None:
55
+ if key in self._contents:
56
+ self._contents.move_to_end(key)
57
+ else:
58
+ self.prune()
59
+
60
+ self._contents[key] = _PpImageCacheEntry(pp_image=pp_image)
61
+
62
+ def prune(self) -> None:
63
+ cutoff = datetime.now(tz=timezone.utc) - _MAX_AGE_OF_PP_IMAGE
64
+
65
+ while self._contents and next(iter(self._contents.values())).last_accessed_at < cutoff:
66
+ self._contents.popitem(last=False)
67
+
68
+
69
+ _pp_image_cache: _PpImageCache
70
+
71
+
72
+ def _worker_job_interact(
73
+ cache_key: object, image: PIL.Image.Image, prompts: cvataa.InteractionPrompts
74
+ ) -> list[cvataa.InteractionResultShape]:
75
+ current_function = cast(cvataa.InteractionFunction, worker_current_function())
76
+ context = _InteractionFunctionContextImpl()
77
+
78
+ if hasattr(current_function, "preprocess_image"):
79
+ pp_image = current_function.preprocess_image(context, image)
80
+ else:
81
+ pp_image = image
82
+
83
+ _pp_image_cache.set(cache_key, pp_image)
84
+
85
+ return current_function.detect(context, pp_image, prompts)
86
+
87
+
88
+ def _worker_job_interact_from_cache(
89
+ cache_key: object, prompts: cvataa.InteractionPrompts
90
+ ) -> list[cvataa.InteractionResultShape] | None:
91
+ current_function = cast(cvataa.InteractionFunction, worker_current_function())
92
+
93
+ pp_image = _pp_image_cache.get(cache_key)
94
+ if pp_image is _PpImageCache._MISSING:
95
+ return None
96
+
97
+ context = _InteractionFunctionContextImpl()
98
+ return current_function.detect(context, pp_image, prompts)
99
+
100
+
101
+ class AgentInteractionFunctionDriver(AgentFunctionDriver[cvataa.InteractionFunctionSpec]):
102
+ FUNCTION_KIND = "interactor"
103
+
104
+ @classmethod
105
+ def init_worker(cls, state_id_generator: object) -> None:
106
+ global _pp_image_cache
107
+ _pp_image_cache = _PpImageCache()
108
+
109
+ @classmethod
110
+ def get_remote_function_fields(cls, spec: cvataa.InteractionFunctionSpec) -> dict[str, Any]:
111
+ fields = {"min_pos_points": spec.min_pos_points}
112
+
113
+ if spec.min_neg_points is not None:
114
+ fields["min_neg_points"] = spec.min_neg_points
115
+
116
+ match spec.min_bounding_boxes:
117
+ case 0:
118
+ fields["startswith_box_optional"] = True
119
+ case 1:
120
+ fields["startswith_box"] = True
121
+ case None:
122
+ pass
123
+ case _:
124
+ assert False, f"Unexpected min_bounding_boxes value: {spec.min_bounding_boxes}"
125
+
126
+ return fields
127
+
128
+ def validate_function_compatibility(self, remote_function: dict) -> None:
129
+ remote_min_pos_points = remote_function["min_pos_points"]
130
+ if remote_min_pos_points < self._function_spec.min_pos_points:
131
+ raise IncompatibleFunctionError(
132
+ "the remote function allows prompts with "
133
+ f"{remote_min_pos_points} positive point(s), but the function object "
134
+ f"requires at least {self._function_spec.min_pos_points}"
135
+ )
136
+
137
+ remote_min_neg_points = remote_function["min_neg_points"]
138
+ if remote_min_neg_points >= 0:
139
+ if self._function_spec.min_neg_points is None:
140
+ raise IncompatibleFunctionError(
141
+ "the remote function allows prompts with negative points, but the "
142
+ "function object does not"
143
+ )
144
+
145
+ if remote_min_neg_points < self._function_spec.min_neg_points:
146
+ raise IncompatibleFunctionError(
147
+ "the remote function allows prompts with "
148
+ f"{remote_min_neg_points} negative point(s), but the function object "
149
+ f"requires at least {self._function_spec.min_neg_points}"
150
+ )
151
+ else:
152
+ if self._function_spec.min_neg_points not in {None, 0}:
153
+ raise IncompatibleFunctionError(
154
+ "the remote function does not allow prompts with negative points, "
155
+ "but the function object requires at least "
156
+ f"{self._function_spec.min_neg_points}"
157
+ )
158
+
159
+ if self._function_spec.min_bounding_boxes is None:
160
+ if remote_function["startswith_box"] or remote_function["startswith_box_optional"]:
161
+ raise IncompatibleFunctionError(
162
+ "the remote function allows prompts with bounding boxes, but the "
163
+ "function object does not"
164
+ )
165
+ elif self._function_spec.min_bounding_boxes == 1:
166
+ if not remote_function["startswith_box"]:
167
+ raise IncompatibleFunctionError(
168
+ "the remote function does not require bounding boxes in prompts, "
169
+ "but the function object does"
170
+ )
171
+
172
+ def _calculate_result_for_interact_ar(
173
+ self, ar_params: dict[str, Any], check_in: object
174
+ ) -> dict[str, Any]:
175
+ pos_points = tuple(map(tuple, ar_params["pos_points"]))
176
+ neg_points = tuple(map(tuple, ar_params["neg_points"]))
177
+ bounding_box = tuple(map(tuple, ar_params["obj_bbox"])) or None
178
+
179
+ if len(pos_points) < self._function_spec.min_pos_points:
180
+ raise BadArError("not enough positive points")
181
+
182
+ if self._function_spec.min_neg_points is None:
183
+ if len(neg_points) > 0:
184
+ raise BadArError("negative points are not supported")
185
+ else:
186
+ if len(neg_points) < self._function_spec.min_neg_points:
187
+ raise BadArError("not enough negative points")
188
+
189
+ if self._function_spec.min_bounding_boxes is None and bounding_box is not None:
190
+ raise BadArError("bounding boxes are not supported")
191
+ elif self._function_spec.min_bounding_boxes == 1 and bounding_box is None:
192
+ raise BadArError("a bounding box is required")
193
+
194
+ prompts = _InteractionPromptsImpl(
195
+ pos_points=pos_points, neg_points=neg_points, bounding_box=bounding_box
196
+ )
197
+
198
+ cache_key = (ar_params["task"], ar_params["frame"], tuple(ar_params.get("roi") or ()))
199
+
200
+ shapes = self._executor.result(
201
+ self._executor.submit(_worker_job_interact_from_cache, cache_key, prompts)
202
+ )
203
+
204
+ if shapes is None:
205
+ sample, _ = self._get_sample_from_ar_params(ar_params)
206
+
207
+ shapes = self._executor.result(
208
+ self._executor.submit(
209
+ _worker_job_interact,
210
+ cache_key,
211
+ self._load_image_for_ar(sample, ar_params),
212
+ prompts,
213
+ )
214
+ )
215
+
216
+ return {"shapes": [attrs.asdict(shape) for shape in shapes]}
217
+
218
+ _CALCULATE_RESULT_PER_AR_TYPE = {"interact": _calculate_result_for_interact_ar}
@@ -146,9 +146,8 @@ class _TrackingFunctionShapeContextImpl(cvataa.TrackingFunctionShapeContext):
146
146
  original_shape_type: str
147
147
 
148
148
 
149
- class AgentTrackingFunctionDriver(AgentFunctionDriver):
149
+ class AgentTrackingFunctionDriver(AgentFunctionDriver[cvataa.TrackingFunctionSpec]):
150
150
  FUNCTION_KIND = "tracker"
151
- _function_spec: cvataa.TrackingFunctionSpec
152
151
 
153
152
  @classmethod
154
153
  def init_worker(cls, state_id_generator: TrackingStateIdGenerator) -> None:
@@ -158,6 +157,10 @@ class AgentTrackingFunctionDriver(AgentFunctionDriver):
158
157
  global _tracking_state_id_generator
159
158
  _tracking_state_id_generator = state_id_generator
160
159
 
160
+ @classmethod
161
+ def get_remote_function_fields(cls, spec: cvataa.TrackingFunctionSpec) -> dict[str, Any]:
162
+ return {"supported_shape_types": sorted(spec.supported_shape_types)}
163
+
161
164
  def validate_function_compatibility(self, remote_function: dict) -> None:
162
165
  remote_supported_shape_types = frozenset(remote_function["supported_shape_types"])
163
166
  unsupported = remote_supported_shape_types - self._function_spec.supported_shape_types
@@ -53,7 +53,10 @@ class CommandGroup:
53
53
 
54
54
  for name, command in self._commands.items():
55
55
  subparser = subparsers.add_parser(name, description=command.description)
56
- subparser.set_defaults(_executor=command.execute)
56
+ subparser.set_defaults(
57
+ _executor=command.execute,
58
+ _needs_client=getattr(command, "needs_client", True),
59
+ )
57
60
  command.configure_parser(subparser)
58
61
 
59
62
  def execute(self) -> None:
@@ -3,12 +3,14 @@
3
3
  # SPDX-License-Identifier: MIT
4
4
 
5
5
  from .command_base import CommandGroup, DeprecatedAlias
6
+ from .commands_config import COMMANDS as COMMANDS_CONFIG
6
7
  from .commands_functions import COMMANDS as COMMANDS_FUNCTIONS
7
8
  from .commands_projects import COMMANDS as COMMANDS_PROJECTS
8
9
  from .commands_tasks import COMMANDS as COMMANDS_TASKS
9
10
 
10
11
  COMMANDS = CommandGroup(description="Perform operations on CVAT resources.")
11
12
 
13
+ COMMANDS.add_command("config", COMMANDS_CONFIG)
12
14
  COMMANDS.add_command("function", COMMANDS_FUNCTIONS)
13
15
  COMMANDS.add_command("project", COMMANDS_PROJECTS)
14
16
  COMMANDS.add_command("task", COMMANDS_TASKS)
@@ -0,0 +1,45 @@
1
+ # Copyright (C) CVAT.ai Corporation
2
+ #
3
+ # SPDX-License-Identifier: MIT
4
+
5
+ from __future__ import annotations
6
+
7
+ import argparse
8
+
9
+ from cvat_sdk.core.auth import AuthStore
10
+
11
+ from .command_base import CommandGroup
12
+ from .common import CriticalError
13
+
14
+ COMMANDS = CommandGroup(description="Manage local CVAT CLI configuration.")
15
+
16
+
17
+ @COMMANDS.command_class("default-server")
18
+ class ConfigDefaultServer:
19
+ needs_client = False
20
+ description = (
21
+ "Print, set, or unset the default server used by the non-profile "
22
+ "credential paths (--auth, CVAT_ACCESS_TOKEN, prompt)."
23
+ )
24
+
25
+ def configure_parser(self, parser: argparse.ArgumentParser) -> None:
26
+ parser.add_argument("server", nargs="?", default=None, help="server URL to remember")
27
+ parser.add_argument("--unset", action="store_true", help="clear the saved default server")
28
+
29
+ def execute(self, args: argparse.Namespace) -> None:
30
+ store = AuthStore()
31
+
32
+ if args.unset and args.server is not None:
33
+ raise CriticalError("Cannot combine a server value with --unset.")
34
+
35
+ if args.unset:
36
+ store.clear_default_server()
37
+ print("Default server cleared.")
38
+ elif args.server is not None:
39
+ if not args.server.strip():
40
+ raise CriticalError("Default server cannot be empty. Use --unset to clear it.")
41
+ store.set_default_server(args.server)
42
+ print(f"Default server is now {args.server!r}.")
43
+ else:
44
+ current = store.get_default_server()
45
+ print(current if current is not None else "(no default server set)")
@@ -8,12 +8,9 @@ import textwrap
8
8
  from collections.abc import Sequence
9
9
  from typing import Any
10
10
 
11
- import cvat_sdk.auto_annotation as cvataa
12
- from cvat_sdk import Client, models
11
+ from cvat_sdk import Client
13
12
 
14
- from .agent import FUNCTION_PROVIDER_NATIVE, run_agent
15
- from .agent_driver_detection import AgentDetectionFunctionDriver
16
- from .agent_driver_tracking import AgentTrackingFunctionDriver
13
+ from .agent import FUNCTION_PROVIDER_NATIVE, get_function_driver_class, run_agent
17
14
  from .command_base import CommandGroup
18
15
  from .common import FunctionLoader, configure_function_implementation_arguments
19
16
 
@@ -40,29 +37,6 @@ class FunctionCreateNative:
40
37
 
41
38
  configure_function_implementation_arguments(parser)
42
39
 
43
- @staticmethod
44
- def _dump_sublabel_spec(
45
- sl_spec: models.SublabelRequest | models.PatchedLabelRequest,
46
- ) -> dict:
47
- result = {
48
- "name": sl_spec.name,
49
- "attributes": [
50
- {
51
- "name": attribute_spec.name,
52
- "input_type": attribute_spec.input_type,
53
- "values": attribute_spec.values,
54
- }
55
- for attribute_spec in getattr(sl_spec, "attributes", [])
56
- ],
57
- }
58
-
59
- if getattr(sl_spec, "type", "any") != "any":
60
- # Add the type conditionally, to stay compatible with older
61
- # CVAT versions when the function doesn't define label types.
62
- result["type"] = sl_spec.type
63
-
64
- return result
65
-
66
40
  def execute(
67
41
  self,
68
42
  client: Client,
@@ -72,32 +46,16 @@ class FunctionCreateNative:
72
46
  function_loader: FunctionLoader,
73
47
  ) -> None:
74
48
  function = function_loader.load()
49
+ driver_class = get_function_driver_class(function.spec)
75
50
 
76
51
  remote_function: dict[str, Any] = {
77
52
  "provider": FUNCTION_PROVIDER_NATIVE,
78
53
  "name": name,
79
54
  "visibility": visibility,
55
+ "kind": driver_class.FUNCTION_KIND,
56
+ **driver_class.get_remote_function_fields(function.spec),
80
57
  }
81
58
 
82
- spec = function.spec
83
-
84
- if isinstance(spec, cvataa.DetectionFunctionSpec):
85
- remote_function["kind"] = AgentDetectionFunctionDriver.FUNCTION_KIND
86
- remote_function["labels_v2"] = []
87
-
88
- for label_spec in spec.labels:
89
- remote_function["labels_v2"].append(self._dump_sublabel_spec(label_spec))
90
-
91
- if sublabels := getattr(label_spec, "sublabels", None):
92
- remote_function["labels_v2"][-1]["sublabels"] = [
93
- self._dump_sublabel_spec(sublabel) for sublabel in sublabels
94
- ]
95
- elif isinstance(spec, cvataa.TrackingFunctionSpec):
96
- remote_function["kind"] = AgentTrackingFunctionDriver.FUNCTION_KIND
97
- remote_function["supported_shape_types"] = sorted(spec.supported_shape_types)
98
- else:
99
- raise cvataa.BadFunctionError(f"Unsupported function spec type: {type(spec).__name__}")
100
-
101
59
  _, response = client.api_client.call_api(
102
60
  "/api/functions",
103
61
  "POST",
@@ -4,27 +4,23 @@
4
4
  # SPDX-License-Identifier: MIT
5
5
 
6
6
  import argparse
7
- import getpass
8
7
  import importlib
9
8
  import importlib.util
10
9
  import logging
11
- import os
12
10
  import sys
13
- import textwrap
14
- from collections.abc import Callable
15
11
  from http.client import HTTPConnection
16
12
  from pathlib import Path
17
13
  from typing import Any
18
14
 
19
15
  import attrs
20
16
  import cvat_sdk.auto_annotation as cvataa
21
- from cvat_sdk.core.client import (
22
- AccessTokenCredentials,
23
- Client,
24
- Config,
25
- Credentials,
26
- PasswordCredentials,
17
+ from cvat_sdk.core.auth import (
18
+ ClientAuthParameters,
19
+ configure_client_auth_arguments,
20
+ make_client_from_cli,
27
21
  )
22
+ from cvat_sdk.core.client import Client
23
+ from cvat_sdk.core.exceptions import AuthStoreError
28
24
 
29
25
  from ..version import VERSION
30
26
  from .parsers import BuildDictAction, parse_function_parameter
@@ -35,82 +31,9 @@ class CriticalError(Exception):
35
31
  pass
36
32
 
37
33
 
38
- CVAT_ACCESS_TOKEN_ENV_VAR = "CVAT_ACCESS_TOKEN" # nosec - a variable name declaration
39
-
40
-
41
- def default_auth_factory() -> Callable[[str], Credentials]:
42
- """
43
- Try to read the CVAT_ACCESS_TOKEN environment variable for a Personal Access Token (PAT).
44
- If there is no value, try using the current user and asking for the password.
45
- """
46
-
47
- token = os.getenv(CVAT_ACCESS_TOKEN_ENV_VAR)
48
- if token is not None:
49
- return lambda _: AccessTokenCredentials(token)
50
-
51
- return get_auth_factory(getpass.getuser())
52
-
53
-
54
- def get_auth_factory(s: str) -> Callable[[str], Credentials]:
55
- """
56
- Parse a USER[:PASS] string and return a callable that takes the server URL
57
- and returns auth credentials for that URL.
58
- The callable will prompt the user for the password if none was initially supplied in the
59
- input string and in the PASS env variable.
60
- """
61
-
62
- user, _, password = s.partition(":")
63
- if not password:
64
- password = os.environ.get("PASS")
65
-
66
- if password:
67
- return lambda _: PasswordCredentials(user, password)
68
- else:
69
- return lambda url: PasswordCredentials(
70
- user, getpass.getpass(f"Password for {user} at {url}: ")
71
- )
72
-
73
-
74
34
  def configure_common_arguments(parser: argparse.ArgumentParser) -> None:
75
35
  parser.add_argument("--version", action="version", version=VERSION)
76
- parser.add_argument(
77
- "--insecure",
78
- action="store_true",
79
- help="Allows to disable SSL certificate check",
80
- )
81
-
82
- parser.add_argument(
83
- "--auth",
84
- type=get_auth_factory,
85
- metavar="USER[:PASS]",
86
- default=default_auth_factory(),
87
- help=textwrap.dedent("""\
88
- User and password to use for authentication;
89
- defaults to the current user and supports the PASS
90
- environment variable or password prompt.
91
- A Personal Access Token (PAT) can be generated on the server
92
- and specified in the {} environment variable instead.
93
- (default user: {}).
94
- """).format(CVAT_ACCESS_TOKEN_ENV_VAR, getpass.getuser()),
95
- )
96
- parser.add_argument(
97
- "--server-host", type=str, default="http://localhost", help="host (default: %(default)s)"
98
- )
99
- parser.add_argument(
100
- "--server-port",
101
- type=int,
102
- default=None,
103
- help="port (default: 80 for http and 443 for https connections)",
104
- )
105
- parser.add_argument(
106
- "--organization",
107
- "--org",
108
- metavar="SLUG",
109
- help="""short name (slug) of the organization
110
- to use when listing or creating resources;
111
- set to blank string to use the personal workspace
112
- (default: list all accessible objects, create in personal workspace)""",
113
- )
36
+ configure_client_auth_arguments(parser)
114
37
  parser.add_argument(
115
38
  "--debug",
116
39
  action="store_const",
@@ -119,6 +42,7 @@ def configure_common_arguments(parser: argparse.ArgumentParser) -> None:
119
42
  default=logging.INFO,
120
43
  help="show debug output",
121
44
  )
45
+ parser.set_defaults(_needs_client=True)
122
46
 
123
47
 
124
48
  def configure_logger(logger: logging.Logger, parsed_args: argparse.Namespace) -> None:
@@ -135,24 +59,16 @@ def configure_logger(logger: logging.Logger, parsed_args: argparse.Namespace) ->
135
59
 
136
60
 
137
61
  def build_client(parsed_args: argparse.Namespace, logger: logging.Logger) -> Client:
138
- config = Config(verify_ssl=not popattr(parsed_args, "insecure"))
139
-
140
- url = popattr(parsed_args, "server_host")
141
- if server_port := popattr(parsed_args, "server_port"):
142
- url += f":{server_port}"
62
+ auth_args = ClientAuthParameters.from_namespace(parsed_args)
63
+ for field in attrs.fields(ClientAuthParameters):
64
+ popattr(parsed_args, field.name)
143
65
 
144
- client = Client(
145
- url=url,
146
- logger=logger,
147
- config=config,
148
- check_server_version=False, # version is checked after auth to support versions < 2.3
149
- )
66
+ try:
67
+ client = make_client_from_cli(auth_args, logger=logger)
68
+ except AuthStoreError as e:
69
+ raise CriticalError(str(e)) from e
150
70
 
151
- client.login(popattr(parsed_args, "auth")(client.api_client.configuration.host))
152
71
  client.check_server_version(fail_if_unsupported=False)
153
-
154
- client.organization_slug = popattr(parsed_args, "organization")
155
-
156
72
  return client
157
73
 
158
74
 
@@ -0,0 +1 @@
1
+ VERSION = "2.71.0"
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: cvat-cli
3
- Version: 2.69.0
3
+ Version: 2.71.0
4
4
  Summary: Command-line client for CVAT
5
5
  Author-email: "CVAT.ai Corporation" <support@cvat.ai>
6
6
  License-Expression: MIT
@@ -9,7 +9,7 @@ Classifier: Programming Language :: Python :: 3
9
9
  Classifier: Operating System :: OS Independent
10
10
  Requires-Python: >=3.10
11
11
  Description-Content-Type: text/markdown
12
- Requires-Dist: cvat-sdk==2.69.0
12
+ Requires-Dist: cvat-sdk==2.71.0
13
13
  Requires-Dist: attrs>=24.2.0
14
14
  Requires-Dist: Pillow>=10.3.0
15
15
 
@@ -13,9 +13,11 @@ src/cvat_cli/_internal/__init__.py
13
13
  src/cvat_cli/_internal/agent.py
14
14
  src/cvat_cli/_internal/agent_driver.py
15
15
  src/cvat_cli/_internal/agent_driver_detection.py
16
+ src/cvat_cli/_internal/agent_driver_interaction.py
16
17
  src/cvat_cli/_internal/agent_driver_tracking.py
17
18
  src/cvat_cli/_internal/command_base.py
18
19
  src/cvat_cli/_internal/commands_all.py
20
+ src/cvat_cli/_internal/commands_config.py
19
21
  src/cvat_cli/_internal/commands_functions.py
20
22
  src/cvat_cli/_internal/commands_projects.py
21
23
  src/cvat_cli/_internal/commands_tasks.py
@@ -1,3 +1,3 @@
1
- cvat-sdk==2.69.0
1
+ cvat-sdk==2.71.0
2
2
  attrs>=24.2.0
3
3
  Pillow>=10.3.0
@@ -1 +0,0 @@
1
- VERSION = "2.69.0"
File without changes
File without changes