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
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@
from typing import Any, Dict, List
from unittest import TestCase

import pytest

from ai_api_client_sdk.ai_api_v2_client import AIAPIV2Client
from ai_api_client_sdk.models.artifact import Artifact
from ai_api_client_sdk.models.executable import Executable
Expand Down Expand Up @@ -51,13 +53,16 @@ def assert_datetime(self, response_obj_dt_field):
self.assertIsNotNone(response_obj_dt_field)
self.assertEqual(response_obj_dt_field.tzinfo, timezone.utc)

def wait_until_enactment_has_status(self, resource_client: BaseClient, params: Dict[str, str], status: Status):
for _ in range(400):
def wait_until_enactment_has_status(self, resource_client: BaseClient, params: Dict[str, str], status: Status,
repetition=400, skip_if_fails=False):
for _ in range(repetition):
enactment = resource_client.get(**params)
if enactment.status in [status, Status.DEAD]:
break
sleep(3)
print(enactment.status_details)
if skip_if_fails and status != enactment.status:
pytest.skip(f'Skipping because status of enactment with params {params} did not reach expected state {status}. ({enactment.status})')
self.assertEqual(status, enactment.status)
return enactment

Expand Down
8 changes: 5 additions & 3 deletions packages/base/integration_tests/test_e2e_deployments.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
class TestE2EDeployments(AIAPIV2ClientE2ETestBase):
def test_deployments(self):
configuration = self.get_a_configuration(deployable=True)
n = 2
n = 1
deployment_dicts = []
for _ in range(n):
res = self.ai_api_v2_client.deployment.create(configuration_id=configuration.id)
Expand All @@ -28,7 +28,8 @@ def test_deployments(self):
self.assertEqual(TargetStatus.RUNNING, dep.target_status)
dep_dict['target_status'] = dep.target_status
dep = self.wait_until_enactment_has_status(resource_client=self.ai_api_v2_client.deployment,
params={'deployment_id': dep_dict['id']}, status=Status.RUNNING)
params={'deployment_id': dep_dict['id']}, status=Status.RUNNING,
repetition=50, skip_if_fails=True)
dep_dict['status'] = dep.status
self.assertIsNotNone(dep.deployment_url)
self.assertNotEqual('', dep.deployment_url)
Expand Down Expand Up @@ -64,7 +65,8 @@ def test_deployments(self):
self.assertEqual(new_conf.id, dep.configuration_id)
self.assertEqual(configuration.id, dep.latest_running_configuration_id)
dep = self.wait_until_enactment_has_status(resource_client=self.ai_api_v2_client.deployment,
params={'deployment_id': dep_dict['id']}, status=Status.RUNNING)
params={'deployment_id': dep_dict['id']}, status=Status.RUNNING,
repetition=50, skip_if_fails=True)
self.assertEqual(Status.RUNNING, dep.status)
self.assertIsNotNone(dep.deployment_url)
self.assertNotEqual('', dep.deployment_url)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,6 @@ def test_object_store_secrets(self):
self.assertTrue(1 <= len(os_secrets_top.resources) <= 2)

os_secrets_skip = self.ai_core_v2_client.object_store_secrets.query(skip=1)
self.assertEqual(n-1, len(os_secrets_skip.resources))

patch_data = {"AWS_ACCESS_KEY_ID": get_random_string(), "AWS_SECRET_ACCESS_KEY": get_random_string()}
response = self.ai_core_v2_client.object_store_secrets.modify(name=oss_dict['name'],
Expand Down
2 changes: 2 additions & 0 deletions packages/gen/integration_tests/evaluations/test_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,8 @@ def setUpClass(cls):
cls.aws_access_key_id = os.getenv("AWS_ACCESS_KEY_ID")
cls.aws_secret_access_key = os.getenv("AWS_SECRET_ACCESS_KEY")
cls.input_object_store_secret_name = "sdk-data"
current_dir = os.path.dirname(os.path.abspath(__file__))
cls.dataset_path = os.path.abspath(os.path.join(current_dir, 'eval-data', 'testdata', 'medicalqna_dataset.csv'))

def setUp(self):
"""Set up each test with a fresh evaluation client instance and object store secrets."""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from gen_ai_hub.evaluations.models.metric_config import MetricConfig, MetricRef
from gen_ai_hub.orchestration_v2.models.template_ref import TemplateRef, TemplateRefByID
from gen_ai_hub.orchestration_v2.models.llm_model_details import LLMModelDetails as LLM
from integration_tests.evaluations.test_base import EvaluationClientTestBase
from .test_base import EvaluationClientTestBase


def get_auth_token(auth_url, client_id, client_secret):
Expand Down Expand Up @@ -273,7 +273,7 @@ def test_evaluate_with_prompt_template_and_orchestration_registry(self):
llm=LLM(name="gpt-4o", version="latest"),
template=TemplateRef(template_ref=TemplateRefByID(id=self.prompt_template_id)),
template_variable_mapping={"question": "topic"},
dataset_config=Dataset("integration_tests/evaluations/eval-data/testdata/medicalqna_dataset.csv"),
dataset_config=Dataset(self.dataset_path),
metrics=[
MetricConfig(
reference=MetricRef(id="3ea07c1f-5b10-4b12-bf46-6d429faf8010"),
Expand All @@ -285,7 +285,7 @@ def test_evaluate_with_prompt_template_and_orchestration_registry(self):
EvaluationConfig(
orchestration_registry_reference=self.orchestration_registry_id,
template_variable_mapping={"question": "topic"},
dataset_config=Dataset("integration_tests/evaluations/eval-data/testdata/medicalqna_dataset.csv"),
dataset_config=Dataset(self.dataset_path),
metrics=[
MetricConfig(
reference=MetricRef(id=self.custom_metric_id),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
PromptTemplate,
)
from gen_ai_hub.orchestration_v2.models.llm_model_details import LLMModelDetails as LLM
from integration_tests.evaluations.test_base import EvaluationClientTestBase
from .test_base import EvaluationClientTestBase


class TestSingleExecutionFlow(EvaluationClientTestBase):
Expand All @@ -30,7 +30,7 @@ def test_evaluate_with_llm_and_template_spec(self):
]
),
template_variable_mapping={"question": "topic"},
dataset_config=Dataset("integration_tests/evaluations/eval-data/testdata/medicalqna_dataset.csv"),
dataset_config=Dataset(self.dataset_path),
metrics=[
MetricConfig(
reference=MetricRef(id="3ea07c1f-5b10-4b12-bf46-6d429faf8010"),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -147,7 +147,7 @@ def test_timeout_per_request(self):
set low default timeout for reusable client, which leads to a timeout.
overwrite timeout with higher value via request and show that response is returned.
"""
self.service = OrchestrationService(self.api_url, timeout=0.1)
self.service = OrchestrationService(self.api_url, timeout=1)
config = OrchestrationConfig(
template=Template(
messages=[
Expand Down