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
38 changes: 30 additions & 8 deletions datamint/api/base_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@

_PAGE_LIMIT = 5000


@dataclass
class ApiConfig:
"""Configuration for API client.
Expand All @@ -31,18 +32,26 @@ class ApiConfig:
api_key: Optional API key for authentication.
timeout: Request timeout in seconds.
max_retries: Maximum number of retries for requests.
port: Optional port number for the API server.
"""
server_url: str
api_key: str | None = None
timeout: float = 30.0
max_retries: int = 3
port: int | None = None

@property
def web_app_url(self) -> str:
"""Get the base URL for the web application."""
if self.server_url.startswith('http://localhost:3001'):
base_url = self.server_url

# Add port to base_url if specified
if self.port is not None:
base_url = f"{self.server_url.rstrip('/')}:{self.port}"

if base_url.startswith('http://localhost'):
return 'http://localhost:3000'
if self.server_url.startswith('https://stagingapi.datamint.io'):
if base_url.startswith('https://stagingapi.datamint.io'):
return 'https://staging.datamint.io'
return 'https://app.datamint.io'

Expand All @@ -68,15 +77,29 @@ def __init__(self,
@staticmethod
def _create_client(config: ApiConfig) -> httpx.Client:
"""Create and configure HTTP client with authentication and timeouts.

The client is designed to be long-lived and reused across multiple requests.
It maintains connection pooling for improved performance.
Default limits: max_keepalive_connections=20, max_connections=100
"""
headers = {"apikey": config.api_key} if config.api_key else None
headers = {"apikey": config.api_key, 'Authorization': f"Bearer {config.api_key}"} if config.api_key else None

# Add port to base_url if specified
base_url = config.server_url.rstrip('/').strip()
if config.port is not None:
# if the port is already in the URL, replace it
if ':' in base_url.split('//')[-1]:
parts = base_url.rsplit(':', 1)
# confirm parts[1] is numeric
if parts[1].isdigit():
base_url = f"{parts[0]}:{config.port}"
else:
logger.warning(f"Invalid port detected in server_url: {config.server_url}")
else:
base_url = f"{base_url}:{config.port}"

return httpx.Client(
base_url=config.server_url,
base_url=base_url,
headers=headers,
timeout=config.timeout,
limits=httpx.Limits(
Expand All @@ -88,7 +111,7 @@ def _create_client(config: ApiConfig) -> httpx.Client:

def close(self) -> None:
"""Close the HTTP client and release resources.

Should be called when the API instance is no longer needed.
Only closes the client if it was created by this instance.
"""
Expand Down Expand Up @@ -379,7 +402,7 @@ def _make_request_with_pagination(self,
"""
offset = 0
total_fetched = 0

use_json_pagination = method.upper() == 'POST' and 'json' in kwargs and isinstance(kwargs['json'], dict)

if not use_json_pagination:
Expand Down Expand Up @@ -480,7 +503,6 @@ def convert_format(bytes_array: bytes,
raise ValueError("Could not determine mimetype from content.")
content_io = BytesIO(bytes_array)
if mimetype.endswith('/dicom'):
import pydicom
return pydicom.dcmread(content_io)
elif mimetype.startswith('image/'):
return Image.open(content_io)
Expand Down
40 changes: 30 additions & 10 deletions datamint/api/client.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
from typing import Optional
from .base_api import ApiConfig, BaseApi
from .endpoints import (ProjectsApi, ResourcesApi, AnnotationsApi,
ChannelsApi, UsersApi, DatasetsInfoApi, ModelsApi,
AnnotationSetsApi
ChannelsApi, UsersApi, DatasetsInfoApi,
AnnotationSetsApi, DeployModelApi
)
from .endpoints.models_api import ModelsApi
import datamint.configs
from datamint.exceptions import DatamintException

Expand All @@ -13,7 +13,7 @@ class Api:
DEFAULT_SERVER_URL = 'https://api.datamint.io'
DATAMINT_API_VENV_NAME = datamint.configs.ENV_VARS[datamint.configs.APIKEY_KEY]

_API_MAP : dict[str, type[BaseApi]] = {
_API_MAP: dict[str, type[BaseApi]] = {
'projects': ProjectsApi,
'resources': ResourcesApi,
'annotations': AnnotationsApi,
Expand All @@ -22,11 +22,12 @@ class Api:
'datasets': DatasetsInfoApi,
'models': ModelsApi,
'annotationsets': AnnotationSetsApi,
'deploy': DeployModelApi,
}

def __init__(self,
server_url: str | None = None,
api_key: Optional[str] = None,
api_key: str | None = None,
timeout: float = 60.0, max_retries: int = 2,
check_connection: bool = True) -> None:
"""Initialize the API client.
Expand Down Expand Up @@ -55,8 +56,16 @@ def __init__(self,
timeout=timeout,
max_retries=max_retries
)
self.mlflow_config = ApiConfig(
server_url=server_url,
api_key=api_key,
timeout=timeout,
max_retries=max_retries,
port=5000
)
self._client = None
self._endpoints = {}
self._mlclient = None
self._endpoints: dict[str, BaseApi] = {}
if check_connection:
self.check_connection()

Expand All @@ -67,12 +76,18 @@ def check_connection(self):
raise DatamintException("Error connecting to the Datamint API." +
f" Please check your api_key and/or other configurations.") from e

def _get_endpoint(self, name: str):
if self._client is None:
self._client = BaseApi._create_client(self.config)
def _get_endpoint(self, name: str, is_mlflow: bool = False):
if is_mlflow:
if self._mlclient is None:
self._mlclient = BaseApi._create_client(self.mlflow_config)
client = self._mlclient
else:
if self._client is None:
self._client = BaseApi._create_client(self.config)
client = self._client
if name not in self._endpoints:
api_class = self._API_MAP[name]
endpoint = api_class(self.config, self._client)
endpoint = api_class(self.config, client)
# Inject this API instance into the endpoint so it can inject into entities
endpoint._api_instance = self
self._endpoints[name] = endpoint
Expand Down Expand Up @@ -110,3 +125,8 @@ def models(self) -> ModelsApi:
@property
def annotationsets(self) -> AnnotationSetsApi:
return self._get_endpoint('annotationsets')

@property
def deploy(self) -> DeployModelApi:
"""Access deployment management endpoints."""
return self._get_endpoint('deploy', is_mlflow=True)
4 changes: 2 additions & 2 deletions datamint/api/endpoints/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,8 @@
from .resources_api import ResourcesApi
from .users_api import UsersApi
from .datasetsinfo_api import DatasetsInfoApi
from .models_api import ModelsApi
from .annotationsets_api import AnnotationSetsApi
from .deploy_model_api import DeployModelApi

__all__ = [
'AnnotationsApi',
Expand All @@ -16,6 +16,6 @@
'ResourcesApi',
'UsersApi',
'DatasetsInfoApi',
'ModelsApi',
'AnnotationSetsApi',
'DeployModelApi',
]
6 changes: 4 additions & 2 deletions datamint/api/endpoints/annotations_api.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
from typing import Any, Sequence, Literal, BinaryIO, Generator, IO
from typing import Literal, BinaryIO, IO
from collections.abc import Sequence, Generator
import httpx
from datetime import date
import logging
from ..entity_base_api import ApiConfig, CreatableEntityApi, DeletableEntityApi
from .models_api import ModelsApi
from datamint.entities.annotations.annotation import Annotation
from datamint.entities.resource import Resource
from datamint.api.dto import AnnotationType, CreateAnnotationDto, LineGeometry, BoxGeometry, CoordinateSystem, Geometry
Expand Down Expand Up @@ -42,6 +42,8 @@ def __init__(self,
client: Optional HTTP client instance. If None, a new one will be created.
"""
from .resources_api import ResourcesApi
from .models_api import ModelsApi

super().__init__(config, Annotation, 'annotations', client)
self._models_api = ModelsApi(config, client=client) if models_api is None else models_api
self._resources_api = ResourcesApi(
Expand Down
78 changes: 78 additions & 0 deletions datamint/api/endpoints/deploy_model_api.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
import httpx
from ..entity_base_api import EntityBaseApi, ApiConfig
from datamint.entities.deployjob import DeployJob

class DeployModelApi(EntityBaseApi[DeployJob]):
"""API handler for model deployment endpoints."""

def __init__(self,
config: ApiConfig,
client: httpx.Client | None = None) -> None:
super().__init__(config, DeployJob, 'datamint/api/v1/deploy-model', client)

def get_by_id(self, entity_id: str) -> DeployJob:
"""Get deployment job status by ID."""
response = self._make_request('GET', f'/{self.endpoint_base}/status/{entity_id}')
data = response.json()
if 'job_id' in data:
data['id'] = data.pop('job_id')
return self._init_entity_obj(**data)

def start(self,
model_name: str,
model_version: int | None = None,
model_alias: str | None = None,
image_name: str | None = None,
with_gpu: bool = False,
convert_to_onnx: bool = False,
input_shape: list[int] | None = None) -> DeployJob:
"""Start a new deployment job."""
payload = {
"model_name": model_name,
"model_version": model_version,
"model_alias": model_alias,
"image_name": image_name,
"with_gpu": with_gpu,
"convert_to_onnx": convert_to_onnx,
"input_shape": input_shape
}
# Remove None values
payload = {k: v for k, v in payload.items() if v is not None}

response = self._make_request('POST', f'/{self.endpoint_base}/start', json=payload)
data = response.json()
return self.get_by_id(data['job_id'])

def cancel(self, job: str | DeployJob) -> bool:
"""Cancel a deployment job."""
job_id = self._entid(job)
response = self._make_request('POST', f'/{self.endpoint_base}/cancel/{job_id}')
return response.json().get('success', False)

def list_active_jobs(self) -> dict:
"""List active deployment jobs count."""
response = self._make_request('GET', f'/{self.endpoint_base}/jobs')
return response.json()

def list_images(self, model_name: str | None = None) -> list[dict]:
"""List deployed model images."""
params = {}
if model_name:
params['model_name'] = model_name
response = self._make_request('GET', f'/{self.endpoint_base}/images', params=params)
return response.json()

def remove_image(self, model_name: str, tag: str | None = None) -> dict:
"""Remove a deployed model image."""
params = {}
if tag:
params['tag'] = tag
response = self._make_request('DELETE', f'/{self.endpoint_base}/image/{model_name}', params=params)
return response.json()

def image_exists(self, model_name: str, tag: str = "champion") -> bool:
"""Check if a model image exists."""
params = {'tag': tag}
response = self._make_request('GET', f'/{self.endpoint_base}/image/{model_name}/exists', params=params)
return response.json().get('exists', False)

1 change: 1 addition & 0 deletions datamint/api/endpoints/models_api.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
"""Deprecated: Use MLFlow API instead."""
from typing import Sequence
from ..entity_base_api import ApiConfig, BaseApi
import httpx
Expand Down
3 changes: 2 additions & 1 deletion datamint/api/entity_base_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,8 @@ class EntityBaseApi(BaseApi, Generic[T]):
def __init__(self, config: ApiConfig,
entity_class: Type[T],
endpoint_base: str,
client: httpx.Client | None = None) -> None:
client: httpx.Client | None = None
) -> None:
"""Initialize the entity API handler.

Args:
Expand Down
18 changes: 18 additions & 0 deletions datamint/entities/deployjob.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
from datamint.entities.base_entity import BaseEntity


class DeployJob(BaseEntity):
id: str
status: str
model_name: str
model_version: int | None = None
model_alias: str | None = None
image_name: str | None = None
image_tag: str | None = None
error_message: str | None = None
progress_percentage: int = 0
current_step: str | None = None
with_gpu: bool = False
recent_logs: list[str] | None = None
started_at: str | None = None
completed_at: str | None = None
34 changes: 32 additions & 2 deletions datamint/mlflow/tracking/datamint_store.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,9 @@
from mlflow.store.tracking.rest_store import RestStore
from mlflow.exceptions import MlflowException
from mlflow.utils.proto_json_utils import message_to_json
from functools import partial
import json
from typing_extensions import override


class DatamintStore(RestStore):
Expand All @@ -14,7 +17,7 @@ def __init__(self, store_uri: str, artifact_uri=None, force_valid=True):
from datamint.mlflow.env_utils import setup_mlflow_environment
from mlflow.utils.credentials import get_default_host_creds
setup_mlflow_environment()

if store_uri.startswith('datamint://') or 'datamint.io' in store_uri or force_valid:
self.invalid = False
else:
Expand All @@ -26,7 +29,6 @@ def __init__(self, store_uri: str, artifact_uri=None, force_valid=True):

def create_experiment(self, name, artifact_location=None, tags=None, project_id: str | None = None) -> str:
from mlflow.protos.service_pb2 import CreateExperiment
from mlflow.utils.proto_json_utils import message_to_json
from datamint.mlflow.tracking.fluent import get_active_project_id

if self.invalid:
Expand All @@ -44,3 +46,31 @@ def create_experiment(self, name, artifact_location=None, tags=None, project_id:

response_proto = self._call_endpoint(CreateExperiment, req_body)
return response_proto.experiment_id

@override
def get_experiment_by_name(self, experiment_name, project_id: str | None = None):
from datamint.mlflow.tracking.fluent import get_active_project_id
from mlflow.protos.service_pb2 import GetExperimentByName
from mlflow.entities import Experiment
from mlflow.protos import databricks_pb2

if self.invalid:
return super().get_experiment_by_name(experiment_name)
if project_id is None:
project_id = get_active_project_id()
try:
req_body = message_to_json(GetExperimentByName(experiment_name=experiment_name))
if project_id:
body = json.loads(req_body)
body["project_id"] = project_id
req_body = json.dumps(body)

response_proto = self._call_endpoint(GetExperimentByName, req_body)
return Experiment.from_proto(response_proto.experiment)
except MlflowException as e:
if e.error_code == databricks_pb2.ErrorCode.Name(
databricks_pb2.RESOURCE_DOES_NOT_EXIST
):
return None
else:
raise
Loading
Loading