diff --git a/python/fi/prompt/client.py b/python/fi/prompt/client.py index 9fddb4c..a3c0da4 100644 --- a/python/fi/prompt/client.py +++ b/python/fi/prompt/client.py @@ -69,7 +69,7 @@ def _parse_success(cls, response) -> Dict: prompt_config_raw = pc cfg_src = (prompt_config_raw or {}).get("configuration", {}) cfg = { - "model_name": cfg_src.get("model_name") or cfg_src.get("model"), + "model_name": cfg_src.get("model") or "unavailable", "temperature": cfg_src.get("temperature"), "frequency_penalty": cfg_src.get("frequency_penalty"), "presence_penalty": cfg_src.get("presence_penalty"), @@ -80,7 +80,7 @@ def _parse_success(cls, response) -> Dict: "tools": cfg_src.get("tools"), } model_config = ModelConfig( - model_name=cfg["model_name"] or "unavailable", + model_name=cfg["model_name"], temperature=cfg["temperature"] if cfg["temperature"] is not None else 0, frequency_penalty=cfg["frequency_penalty"] if cfg["frequency_penalty"] is not None else 0, presence_penalty=cfg["presence_penalty"] if cfg["presence_penalty"] is not None else 0, @@ -130,7 +130,13 @@ def _handle_error(cls, response) -> None: if response.status_code == 400: try: detail = response.json() - error_code = detail.get("error_code") if isinstance(detail, dict) else None + # Backend returns `code` (snake_case-style single word); + # accept `errorCode` as a legacy alternative. + error_code = ( + detail.get("code") + if isinstance(detail, dict) + else None + ) except Exception: error_code = None @@ -162,7 +168,7 @@ def _dict_to_prompt_template(item: Dict) -> PromptTemplate: pc = prompt_config_raw[0] if isinstance(prompt_config_raw, list) else prompt_config_raw cfg_raw = pc.get("configuration", {}) cfg = { - "model_name": cfg_raw.get("model_name") or cfg_raw.get("model"), + "model_name": cfg_raw.get("model") or "unavailable", "temperature": cfg_raw.get("temperature"), "frequency_penalty": cfg_raw.get("frequency_penalty"), "presence_penalty": cfg_raw.get("presence_penalty"), @@ -173,7 +179,7 @@ def _dict_to_prompt_template(item: Dict) -> PromptTemplate: "tools": cfg_raw.get("tools"), } model_config = ModelConfig( - model_name=cfg["model_name"] or "unavailable", + model_name=cfg["model_name"], temperature=cfg["temperature"] if cfg["temperature"] is not None else 0, frequency_penalty=cfg["frequency_penalty"] if cfg["frequency_penalty"] is not None else 0, presence_penalty=cfg["presence_penalty"] if cfg["presence_penalty"] is not None else 0, @@ -243,6 +249,22 @@ def __init__( fi_base_url: Optional[str] = None, **kwargs, ): + """Initialize the Prompt client. + + If ``template`` has no ``id`` but has a ``name``, the SDK will attempt + to fetch the corresponding template from the backend. This supports + two workflows: + + 1. Existing template — pass a ``PromptTemplate(name=...)`` and the + constructor will populate ``id``/``version`` from the backend. + 2. New template — pass a ``PromptTemplate(name=..., messages=...)`` + for a name that doesn't yet exist; the fetch will fail softly + (logged warning) and the user-provided template is retained with + ``id=None`` so ``create()`` can register it. + + For explicit retrieval use ``Prompt.get_template_by_name()`` which + raises ``TemplateNotFound`` on miss instead of falling back. + """ super().__init__( fi_api_key=fi_api_key, fi_secret_key=fi_secret_key, @@ -252,6 +274,7 @@ def __init__( # Label requested during draft create; will be assigned on commit self._pending_label: Optional[str] = None + self._last_generation_id: Optional[str] = None if template and not template.id: try: @@ -265,7 +288,16 @@ def __init__( self.template = template def generate(self, requirements: str) -> "Prompt": - """Generate a prompt and return self for chaining""" + """Submit a prompt-generation job to the backend (asynchronous). + + The backend queues the generation and returns a ``generation_id``. + The result is **not** available synchronously — there is currently no + public endpoint in the backend to poll for a generation job's output by + ``generation_id``. The generated prompt is surfaced through the + FutureAGI UI / job queue rather than through this SDK. + + Use ``last_generation_id`` to retrieve the id for logging / correlation. + """ if not self.template: raise ValueError("No template configured") response = self.request( @@ -274,13 +306,20 @@ def generate(self, requirements: str) -> "Prompt": url=self._base_url + "/" + Routes.generate_prompt.value, json={"statement": requirements}, ), - response_handler=PromptResponseHandler, + response_handler=SimpleJsonResponseHandler, + ) + result = response.get("result", response) if isinstance(response, dict) else response + self._last_generation_id = ( + result.get("generation_id") if isinstance(result, dict) else None ) - self.template.messages[-1].content = response["result"]["prompt"] return self def improve(self, requirements: str) -> "Prompt": - """Improve prompt and return self for chaining""" + """Submit a prompt-improvement job to the backend (asynchronous). + + The backend queues the improvement and returns a ``generation_id``. + See ``generate()`` for notes on async result retrieval. + """ if not self.template: raise ValueError("No template configured") @@ -297,11 +336,23 @@ def improve(self, requirements: str) -> "Prompt": "improvement_requirements": requirements, }, ), - response_handler=PromptResponseHandler, + response_handler=SimpleJsonResponseHandler, + ) + result = ( + improved_response.get("result", improved_response) + if isinstance(improved_response, dict) + else improved_response + ) + self._last_generation_id = ( + result.get("generation_id") if isinstance(result, dict) else None ) - self.template.messages[-1].content = improved_response["result"]["prompt"] return self + @property + def last_generation_id(self) -> Optional[str]: + """Return the generation_id from the most recent generate()/improve() call.""" + return self._last_generation_id + def create(self, *, label: Optional[str] = None) -> "Prompt": """Create a draft prompt template and return self for chaining. @@ -426,6 +477,14 @@ def delete(self) -> bool: if not self.template or not self.template.id: raise ValueError("Template ID missing; cannot delete.") + # Invalidate cache for this template before deleting so subsequent + # lookups don't return stale entries. + if self.template.name: + try: + prompt_cache.invalidate(self.template.name) + except Exception: + logger.warning("prompt_cache.invalidate failed during delete()", exc_info=True) + self.request( config=RequestConfig( method=HttpMethod.DELETE, @@ -461,7 +520,7 @@ def delete_template_by_name( tmpl: PromptTemplate = client.request( config=RequestConfig( method=HttpMethod.GET, - url=client._base_url + "/" + Routes.prompt_label_get_by_name.value, + url=client._base_url + "/" + Routes.get_template_by_name.value, params={"name": name}, ), response_handler=PromptResponseHandler, @@ -476,6 +535,10 @@ def delete_template_by_name( ), response_handler=None, ) + try: + prompt_cache.invalidate(name) + except Exception: + logger.warning("prompt_cache.invalidate failed during delete_template_by_name()", exc_info=True) return True finally: client.close() @@ -486,7 +549,7 @@ def _fetch_template_by_name(self, name: str) -> PromptTemplate: response = self.request( config=RequestConfig( method=HttpMethod.GET, - url=self._base_url + "/" + Routes.prompt_label_get_by_name.value, + url=self._base_url + "/" + Routes.get_template_by_name.value, params={"name": name}, ), response_handler=PromptResponseHandler, @@ -518,9 +581,9 @@ def _fetch_template_version_history(self): def list_template_versions(self): """Return full version history as provided by the backend. - Each element in the returned list is the raw JSON entry that includes - at least these keys: ``template_version``, ``is_draft`` and - ``created_at``. + Each element in the returned list is the raw JSON entry. The backend + returns snake_case keys (``template_version``, ``is_draft``, + ``created_at``); callers that need camelCase should normalize. """ return self._fetch_template_version_history() @@ -532,7 +595,8 @@ def _current_version_is_draft(self) -> bool: """Check backend state to know if the current version is still draft.""" history = self._fetch_template_version_history() for entry in history: - if entry.get("template_version") == self.template.version: + entry_version = entry.get("template_version") + if entry_version == self.template.version: return bool(entry.get("is_draft")) # If not found assume draft (conservative) return True diff --git a/python/fi/prompt/label_management.py b/python/fi/prompt/label_management.py index 788e779..374f108 100644 --- a/python/fi/prompt/label_management.py +++ b/python/fi/prompt/label_management.py @@ -268,7 +268,8 @@ def _assign_label_to_template_version_by_names(self, template_name: str, version history = history_resp.json().get("results", []) matched = None for entry in history: - if str(entry.get("template_version")) == version: + entry_version = entry.get("template_version") + if str(entry_version) == version: matched = entry break if not matched: @@ -328,7 +329,8 @@ def _remove_label_from_template_version_by_names(self, template_name: str, versi history = history_resp.json().get("results", []) version_id = None for entry in history: - if str(entry.get("template_version")) == version: + entry_version = entry.get("template_version") + if str(entry_version) == version: for key in ("id", "version_id", "execution_id"): if entry.get(key): version_id = str(entry.get(key)) @@ -357,8 +359,8 @@ def _get_version_id_by_name(self, version_name: str) -> Optional[str]: """Lookup internal version_id by version name via history endpoint.""" history = self._fetch_template_version_history() for entry in history: - if str(entry.get("template_version")) == version_name: - # Try common id keys + entry_version = entry.get("template_version") + if str(entry_version) == version_name: for key in ("id", "version_id", "execution_id"): if entry.get(key): return str(entry[key]) diff --git a/python/tests/test_prompt_fixes.py b/python/tests/test_prompt_fixes.py new file mode 100644 index 0000000..4d6f0fc --- /dev/null +++ b/python/tests/test_prompt_fixes.py @@ -0,0 +1,522 @@ +"""Unit tests for Prompt module fixes from PR #27. + +Covers four fix categories: + 1. Endpoint routing – get_template_by_name and delete_template_by_name + use Routes.get_template_by_name instead of Routes.prompt_label_get_by_name. + 2. generate() / improve() – use SimpleJsonResponseHandler, safe + response unpacking, expose last_generation_id. + 3. Cache invalidation – prompt_cache.invalidate called on delete. + 4. snake_case key handling – model key, error code, version extraction. +""" + +import json +import uuid +from unittest.mock import MagicMock, patch, PropertyMock + +import pytest + +from fi.prompt.client import Prompt, PromptResponseHandler, SimpleJsonResponseHandler +from fi.prompt.cache import prompt_cache +from fi.prompt.types import PromptTemplate, ModelConfig +from fi.utils.errors import TemplateAlreadyExists, TemplateNotFound +from fi.utils.routes import Routes + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _mock_response(data, status_code=200): + resp = MagicMock() + resp.ok = 200 <= status_code < 300 + resp.status_code = status_code + resp.json.return_value = data + resp.text = json.dumps(data) if not isinstance(data, str) else data + resp.url = "http://test/api/" + resp.request = MagicMock() + resp.request.method = "GET" + resp.request.url = "http://test/api/" + return resp + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +@pytest.fixture +def client(): + """Create a Prompt client with a known template (id present → no init fetch).""" + with patch.dict("os.environ", {"FI_API_KEY": "test-key", "FI_SECRET_KEY": "test-secret"}): + tpl = PromptTemplate( + id=uuid.UUID("00000000-0000-0000-0000-000000000001"), + name="test-template", + messages=[{"role": "user", "content": "Hello"}], + ) + p = Prompt(template=tpl) + return p + + +@pytest.fixture +def mock_request(client): + with patch.object(client, "request") as m: + yield m + + +def _config(mock_request, call_index=-1): + """Extract RequestConfig from a mock call (handles keyword-style calls).""" + call = mock_request.call_args if call_index < 0 else mock_request.call_args_list[call_index] + if call.args: + return call.args[0] + return call.kwargs["config"] + + +# =========================================================================== +# 1. Endpoint Routing +# =========================================================================== + +class TestEndpointRouting: + + # ------------------------------------------------------------------ + # _fetch_template_by_name (called from __init__ and elsewhere) + # ------------------------------------------------------------------ + + def test_fetch_by_name_uses_get_template_by_name_route(self, client, mock_request): + """_fetch_template_by_name must use Routes.get_template_by_name.""" + mock_request.return_value = PromptTemplate(name="test") + client._fetch_template_by_name("test") + + cfg = _config(mock_request) + assert Routes.get_template_by_name.value in cfg.url + assert cfg.params["name"] == "test" + + def test_fetch_by_name_not_using_label_route(self, client, mock_request): + """Verify prompt_label_get_by_name is NOT used in _fetch_template_by_name.""" + mock_request.return_value = PromptTemplate(name="test") + client._fetch_template_by_name("test") + + cfg = _config(mock_request) + assert Routes.prompt_label_get_by_name.value not in cfg.url + + def test_fetch_by_name_passes_response_handler(self, client, mock_request): + """_fetch_template_by_name must pass PromptResponseHandler.""" + mock_request.return_value = PromptTemplate(name="test") + client._fetch_template_by_name("test") + + call = mock_request.call_args if hasattr(mock_request, 'call_args') else mock_request.call_args_list[-1] + assert call.kwargs.get("response_handler") is PromptResponseHandler + + # ------------------------------------------------------------------ + # delete_template_by_name (classmethod, creates its own client) + # ------------------------------------------------------------------ + + def test_delete_by_name_uses_get_template_by_name_route(self): + """delete_template_by_name lookup step must use Routes.get_template_by_name.""" + tpl = PromptTemplate( + id=uuid.UUID("00000000-0000-0000-0000-000000000001"), + name="test", + ) + with patch("fi.prompt.client.APIKeyAuth.request") as mock_req: + mock_req.side_effect = [tpl, None] + with patch.dict("os.environ", {"FI_API_KEY": "k", "FI_SECRET_KEY": "s"}): + Prompt.delete_template_by_name("test") + + # First call = lookup + cfg = _config(mock_req, 0) + assert Routes.get_template_by_name.value in cfg.url + assert cfg.params["name"] == "test" + + def test_delete_by_name_not_using_label_route(self): + """delete_template_by_name must NOT use prompt_label_get_by_name.""" + tpl = PromptTemplate( + id=uuid.UUID("00000000-0000-0000-0000-000000000001"), + name="test", + ) + with patch("fi.prompt.client.APIKeyAuth.request") as mock_req: + mock_req.side_effect = [tpl, None] + with patch.dict("os.environ", {"FI_API_KEY": "k", "FI_SECRET_KEY": "s"}): + Prompt.delete_template_by_name("test") + + cfg = _config(mock_req, 0) + assert Routes.prompt_label_get_by_name.value not in cfg.url + + # ------------------------------------------------------------------ + # get_template_by_name fallback path + # ------------------------------------------------------------------ + + def test_get_template_by_name_fallback_uses_correct_route(self): + """When production label fetch fails, fallback must use get_template_by_name.""" + tpl = PromptTemplate( + id=uuid.UUID("00000000-0000-0000-0000-000000000001"), + name="test", + messages=[{"role": "user", "content": "Hi"}], + ) + with patch("fi.prompt.client.APIKeyAuth.request") as mock_req: + mock_req.side_effect = [TemplateNotFound("test"), tpl] + with patch("fi.prompt.cache.prompt_cache.get", return_value=None): + with patch("fi.prompt.cache.prompt_cache.get_stale", return_value=None): + with patch.dict("os.environ", {"FI_API_KEY": "k", "FI_SECRET_KEY": "s"}): + result = Prompt.get_template_by_name("test") + + # Second call = fallback + cfg = _config(mock_req, 1) + assert Routes.get_template_by_name.value in cfg.url + assert Routes.prompt_label_get_by_name.value not in cfg.url + + def test_get_template_by_name_label_path_uses_label_route(self): + """When a label is explicitly requested, use prompt_label_get_by_name.""" + tpl = PromptTemplate( + id=uuid.UUID("00000000-0000-0000-0000-000000000001"), + name="test", + messages=[{"role": "user", "content": "Hi"}], + ) + with patch("fi.prompt.client.APIKeyAuth.request") as mock_req: + mock_req.return_value = tpl + with patch("fi.prompt.cache.prompt_cache.get", return_value=None): + with patch("fi.prompt.cache.prompt_cache.get_stale", return_value=None): + with patch.dict("os.environ", {"FI_API_KEY": "k", "FI_SECRET_KEY": "s"}): + Prompt.get_template_by_name("test", label="production") + + cfg = _config(mock_req, 0) + assert Routes.prompt_label_get_by_name.value in cfg.url + + +# =========================================================================== +# 2. generate() / improve() +# =========================================================================== + +class TestGenerateImprove: + + def test_generate_uses_simple_json_handler(self, client, mock_request): + """generate() must pass SimpleJsonResponseHandler.""" + mock_request.return_value = {"generation_id": "gen-1"} + client.generate("test") + assert mock_request.call_args.kwargs["response_handler"] is SimpleJsonResponseHandler + + def test_generate_returns_self(self, client, mock_request): + """generate() must return self for chaining.""" + mock_request.return_value = {"generation_id": "gen-1"} + result = client.generate("test") + assert result is client + + def test_generate_sets_last_generation_id(self, client, mock_request): + """generate() must extract generation_id from response.""" + mock_request.return_value = {"result": {"generation_id": "gen-123"}} + client.generate("Make a prompt") + assert client.last_generation_id == "gen-123" + + def test_generate_handles_flat_response(self, client, mock_request): + """generate() must handle response without 'result' wrapper.""" + mock_request.return_value = {"generation_id": "gen-456"} + client.generate("test") + assert client.last_generation_id == "gen-456" + + def test_generate_sets_none_when_no_generation_id(self, client, mock_request): + """generate() must set last_generation_id to None when missing.""" + mock_request.return_value = {"status": "queued"} + client.generate("test") + assert client.last_generation_id is None + + def test_generate_sets_none_when_response_not_dict(self, client, mock_request): + """generate() must handle non-dict response gracefully.""" + mock_request.return_value = "some string" + client.generate("test") + assert client.last_generation_id is None + + def test_generate_raises_without_template(self): + """generate() must raise when no template is configured.""" + with patch.dict("os.environ", {"FI_API_KEY": "k", "FI_SECRET_KEY": "s"}): + p = Prompt() + with pytest.raises(ValueError, match="No template configured"): + p.generate("test") + + def test_improve_uses_simple_json_handler(self, client, mock_request): + """improve() must pass SimpleJsonResponseHandler.""" + mock_request.return_value = {"generation_id": "gen-1"} + client.improve("test") + assert mock_request.call_args.kwargs["response_handler"] is SimpleJsonResponseHandler + + def test_improve_returns_self(self, client, mock_request): + """improve() must return self for chaining.""" + mock_request.return_value = {"generation_id": "gen-1"} + result = client.improve("test") + assert result is client + + def test_improve_sets_last_generation_id(self, client, mock_request): + """improve() must extract generation_id from response.""" + mock_request.return_value = {"result": {"generation_id": "gen-789"}} + client.improve("Make it better") + assert client.last_generation_id == "gen-789" + + def test_improve_handles_flat_response(self, client, mock_request): + """improve() must handle response without 'result' wrapper.""" + mock_request.return_value = {"generation_id": "gen-abc"} + client.improve("test") + assert client.last_generation_id == "gen-abc" + + def test_improve_sets_none_when_no_generation_id(self, client, mock_request): + """improve() must set last_generation_id to None when missing.""" + mock_request.return_value = {"status": "ok"} + client.improve("test") + assert client.last_generation_id is None + + def test_improve_raises_without_template(self): + """improve() must raise when no template is configured.""" + with patch.dict("os.environ", {"FI_API_KEY": "k", "FI_SECRET_KEY": "s"}): + p = Prompt() + with pytest.raises(ValueError, match="No template configured"): + p.improve("test") + + +# =========================================================================== +# 3. Cache Invalidation +# =========================================================================== + +class TestCacheInvalidation: + + def test_delete_invalidates_cache(self, client, mock_request): + """delete() must call prompt_cache.invalidate with template name.""" + mock_request.return_value = None + with patch("fi.prompt.client.prompt_cache.invalidate") as mock_inv: + client.delete() + mock_inv.assert_called_once_with("test-template") + + def test_delete_cache_invalidation_before_http(self, client, mock_request): + """delete() must invalidate cache before the HTTP DELETE call.""" + mock_request.return_value = None + with patch("fi.prompt.client.prompt_cache.invalidate") as mock_inv: + client.delete() + # invalidate must be called before request (order of calls) + mock_inv.assert_called_once() + mock_request.assert_called_once() + + def test_delete_template_by_name_invalidates_cache(self): + """delete_template_by_name() must call prompt_cache.invalidate.""" + tpl = PromptTemplate( + id=uuid.UUID("00000000-0000-0000-0000-000000000001"), + name="test", + ) + with patch("fi.prompt.client.APIKeyAuth.request") as mock_req: + mock_req.side_effect = [tpl, None] + with patch("fi.prompt.client.prompt_cache.invalidate") as mock_inv: + with patch.dict("os.environ", {"FI_API_KEY": "k", "FI_SECRET_KEY": "s"}): + Prompt.delete_template_by_name("test") + mock_inv.assert_called_once_with("test") + + def test_delete_cache_invalidation_graceful_on_failure(self, client, mock_request): + """delete() must not crash when prompt_cache.invalidate raises.""" + mock_request.return_value = None + with patch("fi.prompt.client.prompt_cache.invalidate", side_effect=RuntimeError("fail")): + # Should not raise + result = client.delete() + assert result is True + + def test_delete_template_by_name_graceful_on_cache_failure(self): + """delete_template_by_name() must not crash when cache invalidation raises.""" + tpl = PromptTemplate( + id=uuid.UUID("00000000-0000-0000-0000-000000000001"), + name="test", + ) + with patch("fi.prompt.client.APIKeyAuth.request") as mock_req: + mock_req.side_effect = [tpl, None] + with patch("fi.prompt.client.prompt_cache.invalidate", side_effect=RuntimeError("fail")): + with patch.dict("os.environ", {"FI_API_KEY": "k", "FI_SECRET_KEY": "s"}): + result = Prompt.delete_template_by_name("test") + assert result is True + + +# =========================================================================== +# 4. snake_case Key Handling +# =========================================================================== + +class TestSnakeCaseModelKey: + """Config source reads 'model' key (not 'model_name').""" + + def test_parse_success_reads_model_key(self): + """PromptResponseHandler._parse_success must use cfg_src.get('model').""" + data = { + "result": { + "id": str(uuid.uuid4()), + "name": "test", + "prompt_config": [{ + "configuration": { + "model": "gpt-5", + "temperature": 0.7, + } + }] + } + } + resp = _mock_response(data) + resp.request.method = "GET" + result = PromptResponseHandler._parse_success(resp) + assert isinstance(result, PromptTemplate) + assert result.model_configuration.model_name == "gpt-5" + + def test_parse_success_falls_back_to_unavailable(self): + """When 'model' key is missing, model_name must be 'unavailable'.""" + data = { + "result": { + "id": str(uuid.uuid4()), + "name": "test", + "prompt_config": [{ + "configuration": { + "temperature": 0.7, + } + }] + } + } + resp = _mock_response(data) + resp.request.method = "GET" + result = PromptResponseHandler._parse_success(resp) + assert result.model_configuration.model_name == "unavailable" + + def test_dict_to_prompt_template_reads_model_key(self): + """_dict_to_prompt_template must use cfg_raw.get('model').""" + item = { + "id": str(uuid.uuid4()), + "name": "test", + "prompt_config": [{ + "configuration": { + "model": "claude-sonnet-4", + } + }] + } + result = Prompt._dict_to_prompt_template(item) + assert result.model_configuration.model_name == "claude-sonnet-4" + + def test_dict_to_prompt_template_no_config_falls_back(self): + """When prompt_config is missing entirely, must not crash.""" + item = {"name": "test"} + result = Prompt._dict_to_prompt_template(item) + assert result.model_configuration is not None + # Default model name + assert result.model_configuration.model_name == "gpt-4o-mini" + + def test_dict_to_prompt_template_empty_config(self): + """When configuration is empty, model_name must be 'unavailable'.""" + item = { + "name": "test", + "prompt_config": [{"configuration": {}}] + } + result = Prompt._dict_to_prompt_template(item) + assert result.model_configuration.model_name == "unavailable" + + +class TestSnakeCaseErrorCode: + """Error handler reads 'code' key (not 'error_code').""" + + def test_handle_error_400_reads_code_key(self): + """_handle_error must use detail.get('code').""" + resp = _mock_response( + {"code": "TEMPLATE_ALREADY_EXIST", "name": "my-template"}, + status_code=400, + ) + with pytest.raises(TemplateAlreadyExists, match="my-template"): + PromptResponseHandler._handle_error(resp) + + def test_handle_error_400_unknown_code(self): + """Unknown code must fall through to generic SDKException.""" + resp = _mock_response( + {"code": "SOME_OTHER_ERROR", "message": "Something went wrong"}, + status_code=400, + ) + from fi.utils.errors import SDKException + with pytest.raises(SDKException, match="Something went wrong"): + PromptResponseHandler._handle_error(resp) + + def test_handle_error_404_reads_name_from_query(self): + """404 handler must extract name from query string.""" + resp = _mock_response({}, status_code=404) + resp.request.url = "http://test/api/?name=missing-template" + with pytest.raises(TemplateNotFound, match="missing-template"): + PromptResponseHandler._handle_error(resp) + + def test_handle_error_404_fallback_to_unknown(self): + """404 handler must fallback to 'unknown' when no name in query.""" + resp = _mock_response({}, status_code=404) + resp.request.url = "http://test/api/" + with pytest.raises(TemplateNotFound, match="unknown"): + PromptResponseHandler._handle_error(resp) + + +class TestSnakeCaseVersionExtraction: + """Version lookups use 'template_version' key (extracted to variable).""" + + VERSION_HISTORY = [ + {"template_version": "v1", "is_draft": False, "id": "ver-1"}, + {"template_version": "v2", "is_draft": True, "id": "ver-2"}, + ] + + def test_current_version_is_draft_uses_template_version(self, client): + """_current_version_is_draft must read entry.get('template_version').""" + with patch.object(client, "_fetch_template_version_history") as mock_h: + mock_h.return_value = self.VERSION_HISTORY + + client.template.version = "v1" + assert client._current_version_is_draft() is False + + client.template.version = "v2" + assert client._current_version_is_draft() is True + + def test_current_version_is_draft_defaults_true(self, client): + """When version not found in history, must default to draft (conservative).""" + with patch.object(client, "_fetch_template_version_history") as mock_h: + mock_h.return_value = self.VERSION_HISTORY + client.template.version = "v999" + assert client._current_version_is_draft() is True + + def test_get_version_id_by_name_uses_template_version(self, client): + """_get_version_id_by_name must read entry.get('template_version').""" + with patch.object(client, "_fetch_template_version_history") as mock_h: + mock_h.return_value = self.VERSION_HISTORY + assert client._get_version_id_by_name("v1") == "ver-1" + assert client._get_version_id_by_name("v2") == "ver-2" + + def test_get_version_id_by_name_returns_none_when_missing(self, client): + """_get_version_id_by_name must return None for unknown version.""" + with patch.object(client, "_fetch_template_version_history") as mock_h: + mock_h.return_value = self.VERSION_HISTORY + assert client._get_version_id_by_name("v999") is None + + def _mock_search_result(self, template_id, name): + """Build a mock Response for the template-id resolution step.""" + resp = MagicMock() + resp.url = f"http://test/api/?search={name}" + resp.json.return_value = {"results": [{"id": str(template_id), "name": name}]} + return resp + + def _mock_history_result(self, history): + resp = MagicMock() + resp.url = "http://test/api/history" + resp.json.return_value = {"results": history} + return resp + + def test_assign_label_to_template_version_reads_template_version(self, client): + """_assign_label_to_template_version_by_names must extract template_version.""" + tpl_id = uuid.uuid4() + search_resp = self._mock_search_result(tpl_id, "test") + history_resp = self._mock_history_result(self.VERSION_HISTORY) + assign_resp = MagicMock() + assign_resp.json.return_value = {"status": "success"} + + with patch.object(client, "_get_label_id", return_value="lbl-1"): + with patch.object(client, "request") as mock_req: + mock_req.side_effect = [search_resp, history_resp, assign_resp] + result = client._assign_label_to_template_version_by_names( + "test", "v1", "production" + ) + assert result.json.return_value["status"] == "success" + + def test_assign_label_rejects_draft_version(self, client): + """Assigning label to draft version must raise with 'draft' in message.""" + tpl_id = uuid.uuid4() + search_resp = self._mock_search_result(tpl_id, "test") + history_resp = self._mock_history_result(self.VERSION_HISTORY) + + with patch.object(client, "_get_label_id", return_value="lbl-1"): + with patch.object(client, "request") as mock_req: + mock_req.side_effect = [search_resp, history_resp] + from fi.utils.errors import SDKException + with pytest.raises(SDKException, match="draft"): + client._assign_label_to_template_version_by_names( + "test", "v2", "production" + )