File size: 2,943 Bytes
6f539ba
 
 
 
 
 
 
 
 
 
 
 
d16b23d
6f539ba
 
 
 
 
 
 
 
 
 
d16b23d
 
 
 
 
 
 
 
 
 
 
6f539ba
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d16b23d
6f539ba
 
 
 
 
d16b23d
 
 
 
 
 
6f539ba
 
d16b23d
6f539ba
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
"""Bulk-add dataset collaborators to a Segments.ai dataset from a CSV of usernames and roles.

The CSV is expected to have "username" and "role" columns, e.g.:

    username,role
    john,reviewer
    jane,labeler

Valid roles: labeler, reviewer, manager, admin. Rows with an empty role default to "labeler".
"""

import csv
import json

from segments import SegmentsClient
from segments.exceptions import SegmentsError


VALID_ROLES = ["labeler", "reviewer", "manager", "admin"]
DEFAULT_ROLE = "labeler"


def format_error_message(error: SegmentsError) -> str:
    body, separator, possible_fixes = str(error).partition("\nPossible fixes:")

    try:
        parsed_body = json.loads(body.strip())
    except json.JSONDecodeError:
        parsed_body = None

    if isinstance(parsed_body, dict):
        body = " ".join(str(item) for value in parsed_body.values() for item in (value if isinstance(value, list) else [value]))

    return f"{body.strip()}{separator}{possible_fixes}".strip()


def read_collaborators(csv_path: str) -> tuple[list[tuple[str, str]], list[str]]:
    with open(csv_path, newline="") as csv_file:
        reader = csv.DictReader(csv_file)
        if reader.fieldnames is None or "username" not in reader.fieldnames:
            raise ValueError('CSV must have a "username" column')

        collaborators = []
        skipped_rows = []
        for row_index, row in enumerate(reader, start=2):
            username = row["username"].strip()
            if username == "":
                skipped_rows.append(f"Row {row_index}: skipped, missing username")
                continue

            role = (row.get("role") or "").strip() or DEFAULT_ROLE
            if role not in VALID_ROLES:
                raise ValueError(f"Row {row_index}: invalid role '{role}' for username '{username}'. Must be one of {VALID_ROLES}")

            collaborators.append((username, role))

        return collaborators, skipped_rows


def add_collaborators(client: SegmentsClient, dataset_identifier: str, collaborators: list[tuple[str, str]]) -> tuple[list[str], bool]:
    logs = []
    failed_usernames = []
    usernames_by_error_message: dict[str, list[str]] = {}
    for username, role in collaborators:
        try:
            client.add_dataset_collaborator(dataset_identifier, username, role)
            logs.append(f"Added {username} to {dataset_identifier} as {role}")
        except SegmentsError as error:
             message = format_error_message(error)
             usernames_by_error_message.setdefault(message, []).append(username)
             failed_usernames.append(username)

    for message, usernames in usernames_by_error_message.items():
         logs.append(f"Failed to add {', '.join(usernames)}: {message}\n")

    if failed_usernames:
        logs.append(f"{len(failed_usernames)} of {len(collaborators)} usernames failed: {', '.join(failed_usernames)}")

    return logs, len(failed_usernames) > 0