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
37 changes: 35 additions & 2 deletions packages/examples/cvat/exchange-oracle/debug.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
import datetime
import inspect
import json
import uuid
from collections.abc import Generator
from contextlib import ExitStack, contextmanager
from logging import Logger
Expand All @@ -14,9 +16,10 @@
from src.core.config import Config
from src.db import SessionLocal
from src.services import cloud
from src.services import cvat as cvat_service
from src.services import cvat as cvat_db_service
from src.services.cloud import BucketAccessInfo
from src.utils.logging import format_sequence, get_function_logger
from src.utils.time import utcnow


@contextmanager
Expand Down Expand Up @@ -110,6 +113,8 @@ def _mock_webhook_signature_checking(_: Logger) -> Generator[None, None, None]:
- from reputation oracle -
encoded with Config.localhost.reputation_oracle_address wallet address
or signature "reputation_oracle<number>"

<number> is optional in all cases.
"""

from src.chain.escrow import (
Expand All @@ -133,6 +138,33 @@ def patched_get_available_webhook_types(chain_id, escrow_address):
d[Config.localhost.reputation_oracle_address.lower()] = OracleWebhookTypes.reputation_oracle
return d

from src.services.webhook import inbox as original_inbox

class PatchedInbox:
def __init__(self):
pass

def __getattr__(self, name: str):
return getattr(original_inbox, name)

def create_webhook(
self,
session,
escrow_address,
chain_id,
type: OracleWebhookTypes,
signature=None,
event_type=None,
event_data=None,
event=None,
):
if signature in OracleWebhookTypes:
signature = f"{type.value}-{utcnow().isoformat(sep='T')}-{uuid.uuid4()}"

_orig_params = inspect.signature(original_inbox.create_webhook).parameters
_args = {k: v for k, v in locals().items() if k in _orig_params}
return original_inbox.create_webhook(**_args)

with (
mock.patch("src.schemas.webhook.validate_address", lambda x: x),
mock.patch(
Expand All @@ -143,6 +175,7 @@ def patched_get_available_webhook_types(chain_id, escrow_address):
"src.endpoints.webhook.validate_oracle_webhook_signature",
patched_validate_oracle_webhook_signature,
),
mock.patch("src.services.webhook.inbox", PatchedInbox()),
):
yield

Expand All @@ -165,7 +198,7 @@ def decode_plain_json_token(self, token) -> dict[str, Any]:

if (user_wallet := token_data.get("wallet_address")) and not token_data.get("email"):
with SessionLocal.begin() as session:
user = cvat_service.get_user_by_id(session, user_wallet)
user = cvat_db_service.get_user_by_id(session, user_wallet)
if not user:
raise Exception(f"Could not find user with wallet address '{user_wallet}'")

Expand Down
6 changes: 3 additions & 3 deletions packages/examples/cvat/exchange-oracle/poetry.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

8 changes: 5 additions & 3 deletions packages/examples/cvat/exchange-oracle/src/cvat/api_calls.py
Original file line number Diff line number Diff line change
Expand Up @@ -359,7 +359,7 @@ def put_task_data(
task_id: int,
cloudstorage_id: int,
*,
chunk_size: int,
chunk_size: int | None = None,
filenames: list[str] | None = None,
sort_images: bool | None = None,
validation_params: dict[str, str | float | list[str]] | None = None,
Expand Down Expand Up @@ -404,8 +404,10 @@ def put_task_data(
else models.SortingMethod("predefined")
)

if chunk_size is not None:
kwargs["chunk_size"] = chunk_size

data_request = models.DataRequest(
chunk_size=chunk_size,
cloud_storage_id=cloudstorage_id,
image_quality=Config.cvat_config.image_quality,
use_cache=True,
Expand All @@ -414,7 +416,7 @@ def put_task_data(
**kwargs,
)
try:
(_, response) = api_client.tasks_api.create_data(task_id, data_request=data_request)
api_client.tasks_api.create_data(task_id, data_request=data_request)
return

except exceptions.ApiException as e:
Expand Down
29 changes: 18 additions & 11 deletions packages/examples/cvat/exchange-oracle/src/handlers/job_creation.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,10 +166,6 @@ def _task_segment_size(self) -> int:
def _job_val_frames_count(self) -> int:
return self.manifest.validation.val_size

@property
def _task_chunk_size(self) -> int:
return self._task_segment_size + self._job_val_frames_count

def __enter__(self):
return self

Expand Down Expand Up @@ -434,7 +430,6 @@ def build(self):
cvat_task.id,
cloud_storage.id,
filenames=data_subset,
chunk_size=self._task_chunk_size,
validation_params={
"gt_filenames": gt_filenames, # include whole GT dataset into each task
"gt_frames_per_job_count": self._job_val_frames_count,
Expand Down Expand Up @@ -1614,7 +1609,6 @@ def _create_on_cvat(self):
cvat_task.id,
cvat_cloud_storage.id,
filenames=filenames,
chunk_size=self._task_chunk_size,
validation_params={
"gt_filenames": gt_filenames,
"gt_frames_per_job_count": self._job_val_frames_count,
Expand Down Expand Up @@ -1792,19 +1786,33 @@ def _validate_gt_labels(self):
for node_label in skeleton_label.nodes:
manifest_labels.add((node_label, skeleton_label.name))

if gt_labels - manifest_labels:
if manifest_labels - gt_labels:
raise DatasetValidationError(
"GT labels do not match job labels. Unknown labels: {}".format(
"Could not find GT for labels {}".format(
format_sequence(
[
label_name if not parent_name else f"{parent_name}.{label_name}"
for label_name, parent_name in gt_labels - manifest_labels
for label_name, parent_name in manifest_labels - gt_labels
]
),
)
)

# Reorder labels to match the manifest
# It should not be an issue that there are some extra GT labels - they should
# just be skipped.
if gt_labels - manifest_labels:
self.logger.info(
"Skipping unknown GT labels: {}".format(
format_sequence(
[
label_name if not parent_name else f"{parent_name}.{label_name}"
for label_name, parent_name in gt_labels - manifest_labels
]
)
)
)

# Reorder and filter labels to match the manifest
self._input_gt_dataset.transform(
ProjectLabels, dst_labels=[label.name for label in self.manifest.annotation.labels]
)
Expand Down Expand Up @@ -2942,7 +2950,6 @@ def _task_params_label_key(ts):
cvat_task.id,
cvat_cloud_storage.id,
filenames=point_label_filenames + gt_point_label_filenames,
chunk_size=self._task_chunk_size,
validation_params={
"gt_filenames": gt_point_label_filenames,
"gt_frames_per_job_count": self._job_val_frames_count,
Expand Down
40 changes: 0 additions & 40 deletions packages/examples/cvat/exchange-oracle/src/services/cvat.py
Original file line number Diff line number Diff line change
Expand Up @@ -329,46 +329,6 @@ def create_escrow_validations(session: Session, *, limit: int = 100) -> list[tup
return session.execute(insert_stmt).all()


def get_available_projects(session: Session, *, limit: int = 10) -> list[Project]:
return (
session.query(Project)
.where(
(Project.status == ProjectStatuses.annotation.value)
& Project.jobs.any(
(Job.status == JobStatuses.new)
& ~Job.assignments.any(Assignment.status == AssignmentStatuses.created.value)
)
)
.distinct()
.limit(limit)
.all()
)


def get_projects_by_assignee(
session: Session,
wallet_address: str | None = None,
*,
limit: int = 10,
for_update: bool | ForUpdateParams = False,
) -> list[Project]:
return (
_maybe_for_update(session.query(Project), enable=for_update)
.where(
Project.jobs.any(
Job.assignments.any(
(Assignment.user_wallet_address == wallet_address)
& (Assignment.status == AssignmentStatuses.created)
& (utcnow() < Assignment.expires_at)
)
)
)
.distinct()
.limit(limit)
.all()
)


def update_project_status(session: Session, project_id: str, status: ProjectStatuses) -> None:
upd = update(Project).where(Project.id == project_id).values(status=status.value)
session.execute(upd)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,6 @@
)
from src.db import SessionLocal
from src.models.cvat import Assignment, DataUpload, Image, Job, Project, Task, User
from src.utils.time import utcnow

from tests.utils.db_helper import (
create_project,
Expand Down Expand Up @@ -347,90 +346,6 @@ def test_get_projects_by_status(self):

assert len(projects) == 1

def test_get_available_projects(self):
cvat_id_1 = 456
(cvat_project, cvat_task, cvat_job) = create_project_task_and_job(
self.session, "0x86e83d346041E8806e352681f3F14549C0d2BC67", cvat_id_1
)

projects = cvat_service.get_available_projects(self.session)

assert len(projects) == 1

cvat_id_2 = 457
(cvat_project, cvat_task) = create_project_and_task(
self.session, "0x86e83d346041E8806e352681f3F14549C0d2BC68", cvat_id_2
)

cvat_task_id = cvat_task.cvat_id
cvat_project_id = cvat_project.cvat_id

cvat_service.create_job(
session=self.session,
cvat_id=cvat_id_2,
cvat_task_id=cvat_task_id,
cvat_project_id=cvat_project_id,
status=JobStatuses.in_progress,
start_frame=0,
stop_frame=1,
)

cvat_id_3 = 458
(cvat_project, cvat_task, _) = create_project_task_and_job(
self.session, "0x86e83d346041E8806e352681f3F14549C0d2BC69", cvat_id_3
)

projects = cvat_service.get_available_projects(self.session)
assert len(projects) == 2
assert any(project.cvat_id == cvat_id_1 for project in projects)
assert any(project.cvat_id == cvat_id_3 for project in projects)

def test_get_projects_by_assignee(self):
wallet_address_1 = "0x86e83d346041E8806e352681f3F14549C0d2BC60"
cvat_id_1 = 456

create_project_task_and_job(
self.session, "0x86e83d346041E8806e352681f3F14549C0d2BC67", cvat_id_1
)

user = User(wallet_address=wallet_address_1, cvat_id=cvat_id_1, cvat_email="test@hmt.ai")
self.session.add(user)

cvat_service.create_assignment(
session=self.session,
wallet_address=wallet_address_1,
cvat_job_id=cvat_id_1,
expires_at=datetime.now() + timedelta(days=1),
)

wallet_address_2 = "0x86e83d346041E8806e352681f3F14549C0d2BC61"
cvat_id_2 = 457

create_project_task_and_job(
self.session, "0x86e83d346041E8806e352681f3F14549C0d2BC68", cvat_id_2
)

user = User(wallet_address=wallet_address_2, cvat_id=cvat_id_2, cvat_email="test2@hmt.ai")
self.session.add(user)

cvat_service.create_assignment(
session=self.session,
wallet_address=wallet_address_2,
cvat_job_id=cvat_id_2,
expires_at=utcnow(),
)

projects = cvat_service.get_projects_by_assignee(self.session, wallet_address_1)

assert len(projects) == 1
assert projects[0].cvat_id == cvat_id_1

projects = cvat_service.get_projects_by_assignee(self.session, wallet_address_2)

assert (
len(projects) == 0
) # expired should not be shown, https://github.com/humanprotocol/human-protocol/pull/1879

def test_update_project_status(self):
cvat_id = 1
cvat_cloudstorage_id = 1
Expand Down
Loading