Commit 024c1aaa authored by Open Science Conservation Fund's avatar Open Science Conservation Fund
Browse files

added support for new SDZWA species models (andes, southwest, amazon);...

added support for new SDZWA species models (andes, southwest, amazon); improved cli tool load_aiproviders
parent 0b25e6bf
Loading
Loading
Loading
Loading
+0 −1
Original line number Diff line number Diff line
@@ -23,7 +23,6 @@ class Species(models.Model):
    taxon_rank = models.CharField(
        max_length=20, default=TaxonRank.SPECIES, choices=TaxonRank.CHOICES
    )
    # alternative_name = models.CharField(max_length=100, blank=True, verbose_name=_("Alternative name"))

    class Meta:
        ordering = ["english_name"]
+0 −10
Original line number Diff line number Diff line
@@ -278,7 +278,6 @@ class AIProviderAdmin(PolymorphicParentModelAdmin):
        "crop_image",
        "minimum_confidence",
        "trapper_instance_url",
        "get_api_url",
    )
    inlines = [
        ObservationTypeAILabelMappingInline,
@@ -294,15 +293,6 @@ class AIProviderAdmin(PolymorphicParentModelAdmin):

    actions = ["create_classificator_from_species_mapping"]

    def get_api_url(self, obj):
        """Display api_url for TrapperAIProvider instances"""
        if isinstance(obj, TrapperAIProvider):
            url = obj.api_url
            return mark_safe(f'<a href="{url}" target="_blank">{url}</a>')
        return "-"

    get_api_url.short_description = "API URL"

    @admin.action(description="Create Classificator from AI Provider species mapping")
    def create_classificator_from_species_mapping(self, request, queryset):
        """
+124 −23
Original line number Diff line number Diff line
import getpass
import json
from typing import List
from urllib.parse import urljoin
@@ -64,13 +65,36 @@ class Command(BaseCommand):
            for value_label in species_labels_mapping:
                value = value_label["value"]
                latin_name = value_label["species"]
                common_name = value_label.get("common", None)
                taxon_id = value_label.get("taxon_id", None)
                exclude = value_label.get("exclude", False)

                if exclude:
                    self.stdout.write(
                        self.style.WARNING(
                            f"[WARNING] Skipping label '{latin_name}' (value: {value}) as it is marked to be excluded"
                        )
                    )
                    continue

                # Get species or create a new one
                try:
                species = Species.objects.filter(latin_name=latin_name).first()
                except Species.DoesNotExist:
                if not species:
                    species = Species.objects.create(
                        latin_name=latin_name, taxon_id="anonymus"
                        latin_name=latin_name,
                        english_name=common_name,
                        taxon_id=taxon_id,
                    )
                    self.stdout.write(
                        self.style.NOTICE(
                            f"[INFO] Created new Species: {latin_name} ({common_name})"
                        )
                    )
                else:
                    self.stdout.write(
                        self.style.NOTICE(
                            f"[INFO] Found existing Species: {latin_name} ({common_name})"
                        )
                    )

                species_labels.append(
@@ -125,21 +149,35 @@ class Command(BaseCommand):
        Define positional & optional arguments
        """

        parser.add_argument("trapper_instance_url", type=str, help="Trapper Expert URL")
        parser.add_argument("api_url", type=str, help="TrapperAI Manager URL")
        parser.add_argument("api_auth_login", type=str, help="TrapperAI Manager login")
        parser.add_argument(
            "api_auth_passw", type=str, help="TrapperAI Manager password"
            "--trapper-url",
            type=str,
            help="Trapper Expert URL (optional, inferred from DOMAIN_NAME if not provided)",
        )
        parser.add_argument(
            "--api-url",
            type=str,
            help="Trapper AI Manager URL (required for loading providers)",
        )
        parser.add_argument(
            "--api-login",
            type=str,
            help="Trapper AI Manager login (will prompt interactively if not provided)",
        )
        parser.add_argument(
            "--api-pass",
            type=str,
            help="Trapper AI Manager password (will prompt interactively if not provided)",
        )
        parser.add_argument(
            "--model-codes",
            type=str,
            help="List of model codes to load separated by comma (optional)",
            help="List of model codes to load/display, separated by comma (optional)",
        )
        parser.add_argument(
            "--show-config",
            action="store_true",
            help="Show prediction models config",
            help="Show prediction models config (no credentials required)",
        )
        parser.add_argument(
            "--show-active-models",
@@ -147,25 +185,92 @@ class Command(BaseCommand):
            help="Show active prediction models from TrapperAI Manager",
        )

    def _filter_config_by_model_codes(
        self, config_data: List[dict], model_codes: str
    ) -> List[dict]:
        """
        Filter config data by model codes (comma-separated string)
        """
        if not model_codes:
            return config_data
        codes = [c.strip() for c in model_codes.split(",")]
        return [model for model in config_data if model["code"] in codes]

    def handle(self, *args, **options):
        self.species_count = Species.objects.count()

        with open(settings.AI_PROVIDERS_CONFIG, "r") as file:
            config_data = json.load(file)

            # Show config
            model_codes = options.get("model_codes")

            # Show config (no credentials required)
            if options.get("show_config"):
                self.stdout.write(json.dumps(config_data, indent=4))
                filtered_config = self._filter_config_by_model_codes(
                    config_data, model_codes
                )
                if not filtered_config:
                    self.stdout.write(
                        self.style.WARNING(
                            "[WARNING] No models found for the specified model codes"
                        )
                    )
                else:
                    self.stdout.write(json.dumps(filtered_config, indent=4))
                return

            # Load only specific models if provided
            model_codes = options.get("model_codes")
            # For all other operations, credentials are required
            trapperai_manager_url = options.get("api_url")
            api_auth_login = options.get("api_login")
            api_auth_passw = options.get("api_pass")
            trapper_instance_url = options.get("trapper_url")

            # Infer trapper URL from DOMAIN_NAME if not provided
            if not trapper_instance_url:
                domain_name = getattr(settings, "DOMAIN_NAME", None)
                if domain_name:
                    # if domain contains localhost use http, otherwise https
                    if "localhost" in domain_name or "127.0.0.1" in domain_name:
                        trapper_instance_url = f"http://{domain_name}"
                    else:
                        trapper_instance_url = f"https://{domain_name}"
                    self.stdout.write(
                        self.style.NOTICE(
                            f"[INFO] Inferred Trapper URL: {trapper_instance_url}"
                        )
                    )
                else:
                    self.stdout.write(
                        self.style.ERROR(
                            "[ERROR] Could not infer Trapper URL. Please provide --trapper-url"
                        )
                    )
                    return

            # Validate required --api-url
            if not trapperai_manager_url:
                self.stdout.write(
                    self.style.ERROR("[ERROR] Missing required argument: --api-url")
                )
                return

            # Prompt for credentials interactively if not provided
            if not api_auth_login:
                api_auth_login = input("Trapper AI Manager login: ")
            if not api_auth_passw:
                api_auth_passw = getpass.getpass("Trapper AI Manager password: ")

            if not api_auth_login or not api_auth_passw:
                self.stdout.write(
                    self.style.ERROR("[ERROR] Login and password are required")
                )
                return

            # Filter config by model codes if provided
            if model_codes:
                model_codes = model_codes.split(",")
                config_data = [
                    model for model in config_data if model["code"] in model_codes
                ]
                # If no models found, exit
                config_data = self._filter_config_by_model_codes(
                    config_data, model_codes
                )
                if not config_data:
                    self.stdout.write(
                        self.style.WARNING(
@@ -174,12 +279,8 @@ class Command(BaseCommand):
                    )
                    return

            # Data from positional arguments
            trapperai_manager_url = options.get("api_url")
            api_auth_login = options.get("api_auth_login")
            api_auth_passw = options.get("api_auth_passw")
            ai_provider_auth = {
                "trapper_instance_url": options.get("trapper_instance_url"),
                "trapper_instance_url": trapper_instance_url,
                "api_url": trapperai_manager_url,
                "api_auth_login": api_auth_login,
                "api_auth_passw": api_auth_passw,
+2 −2
Original line number Diff line number Diff line
@@ -294,7 +294,7 @@ class ClassificationProject(models.Model):
    )
    observation_type_confidence_warning_threshold = models.FloatField(
        _("Observation type confidence warning threshold"),
        default=None,
        default=0.7,
        blank=True,
        null=True,
        validators=[MinValueValidator(0.0), MaxValueValidator(1.0)],
@@ -302,7 +302,7 @@ class ClassificationProject(models.Model):
    )
    species_confidence_warning_threshold = models.FloatField(
        _("Species confidence warning threshold"),
        default=None,
        default=0.7,
        blank=True,
        null=True,
        validators=[MinValueValidator(0.0), MaxValueValidator(1.0)],