diff --git a/src/xwhy/explainers/image.py b/src/xwhy/explainers/image.py index dd330239..a86273c5 100644 --- a/src/xwhy/explainers/image.py +++ b/src/xwhy/explainers/image.py @@ -450,6 +450,26 @@ def explain( logger.info("Extracting target probabilities for the predicted class...") y_target = predictions[:, int(class_to_explain)] + # --------------------------------------------------------- + # Distance Validation & Imputation setup: + # Convert distances to numpy array and impute non-finite (inf/NaN) values. + # --------------------------------------------------------- + logger.info("Validating perturbation distances...") + distances_raw = np.array(distances, dtype=float) + + # Filter out non-finite values to determine the maximum valid distance + valid_distances = distances_raw[np.isfinite(distances_raw)] + + # Calculate max_penalty: max valid distance + 1000, or default 1000 if + # all failed + if len(valid_distances) > 0: + max_penalty = np.max(valid_distances) + 1000.0 + else: + max_penalty = 1000.0 + + # Impute infinite/NaN values with the dynamically calculated maximum penalty + distances = np.where(np.isfinite(distances_raw), distances_raw, max_penalty) + if self.config.use_best_surrogate: # type: ignore[union-attr] logger.info("Searching for the optimal surrogate model...") method, score = SurrogateTrainer.find_best( @@ -1378,7 +1398,27 @@ def explain( # Weights: Derived from textual distance (WMD/sims). # --------------------------------------------------------- x_features = np.vstack([np.array(m, dtype=int) for m in binary_masks]) - y_target = image_distances + + # TODO: Modify this maximum distance imputation strategy later. + # Currently using a hardcoded large number (1000.0). Consider updating to + # dynamically calculate the max penalty based on valid distances. + # DO it for all the explainers. + + # Convert image_distances to a numpy array for vectorized imputation + y_target_raw = np.array(image_distances, dtype=float) + + # Filter out infinite values to find the actual maximum valid distance + valid_distances = y_target_raw[np.isfinite(y_target_raw)] + + # Calculate max_penalty: max valid distance + 1000, or just 1000 if all failed + if len(valid_distances) > 0: + max_penalty = np.max(valid_distances) + 1000.0 + else: + max_penalty = 1000.0 + + # Impute infinite values with the dynamically calculated maximum penalty + y_target = np.where(np.isinf(y_target_raw), max_penalty, y_target_raw) + text_distances_array = np.array([d for _, d in wmd_scores]) if self.config.use_best_surrogate: # type: ignore[union-attr] diff --git a/src/xwhy/explainers/llm.py b/src/xwhy/explainers/llm.py index 46d94856..19a82415 100644 --- a/src/xwhy/explainers/llm.py +++ b/src/xwhy/explainers/llm.py @@ -219,12 +219,40 @@ def explain( logger.info("Computing WMD scores...") wmd_distance = WMDDistance() - wmd_scores = wmd_distance.compute_batch( + raw_wmd_scores = wmd_distance.compute_batch( model=self.state.embedding_model, original=original_output, perturbed_texts=perturbed_texts, ) + # --------------------------------------------------------- + # Distance Validation & Imputation setup: + # Convert distances to numpy array and impute non-finite (inf/NaN) values. + # --------------------------------------------------------- + logger.info("Validating perturbation distances...") + distances_raw = np.array([d for _, d in raw_wmd_scores], dtype=float) + + # Filter out non-finite values to determine the maximum valid distance + valid_distances = distances_raw[np.isfinite(distances_raw)] + + # Calculate max_penalty: max valid distance + 1000, or default 1000 if + # all failed + if len(valid_distances) > 0: + max_penalty = np.max(valid_distances) + 1000.0 + else: + max_penalty = 1000.0 + + # Impute infinite/NaN values with the dynamically calculated maximum penalty + distances_array = np.where( + np.isfinite(distances_raw), distances_raw, max_penalty + ) + + # Reconstruct wmd_scores with imputed values for downstream consistency + wmd_scores = [ + (text, float(dist)) + for (text, _), dist in zip(raw_wmd_scores, distances_array, strict=False) + ] + logger.info("Normalizing similarities...") sims = DistanceNormalizer.min_max(scores=wmd_scores) @@ -234,7 +262,6 @@ def explain( x_matrix = np.vstack(masks_as_arrays) y_target = np.array([s for _, s in sims]) - distances_array = np.array([d for _, d in wmd_scores]) if self.config.use_best_surrogate: # type: ignore[union-attr] logger.info( diff --git a/src/xwhy/explainers/tabular.py b/src/xwhy/explainers/tabular.py index 56b3c8d0..a2ec0255 100644 --- a/src/xwhy/explainers/tabular.py +++ b/src/xwhy/explainers/tabular.py @@ -244,6 +244,28 @@ def explain( scaled_distances = distances * cfg.epsilon + # --------------------------------------------------------- + # Distance Validation & Imputation setup: + # Convert distances to numpy array and impute non-finite (inf/NaN) values. + # --------------------------------------------------------- + logger.info("Validating perturbation distances...") + distances_raw = np.array(scaled_distances, dtype=float) + + # Filter out non-finite values to determine the maximum valid distance + valid_distances = distances_raw[np.isfinite(distances_raw)] + + # Calculate max_penalty: max valid distance + 1000, or default 1000 if + # all failed + if len(valid_distances) > 0: + max_penalty = np.max(valid_distances) + 1000.0 + else: + max_penalty = 1000.0 + + # Impute infinite/NaN values with the dynamically calculated maximum penalty + scaled_distances = np.where( + np.isfinite(distances_raw), distances_raw, max_penalty + ) + # 4. Surrogate Training via Framework if cfg.use_best_surrogate: logger.info("Searching for optimal surrogate model...") @@ -296,7 +318,7 @@ def explain( "y_target": y_target, "y_pred": y_pred, "weights": weights, - "distances": distances, + "distances": scaled_distances, "surrogate_method": method, } diff --git a/src/xwhy/explainers/text.py b/src/xwhy/explainers/text.py index bf3943d2..567d4cb5 100644 --- a/src/xwhy/explainers/text.py +++ b/src/xwhy/explainers/text.py @@ -238,14 +238,40 @@ def explain( if self.state.embedding_model is None: raise RuntimeError("Embedding model state is not initialized.") - wmd_scores = wmd_distance.compute_batch( + raw_wmd_scores = wmd_distance.compute_batch( model=self.state.embedding_model, original=instance, perturbed_texts=perturbed_texts, sanitize=True, ) - distances_array = np.array([d for _, d in wmd_scores], dtype=float) + # --------------------------------------------------------- + # Distance Validation & Imputation setup: + # Convert distances to numpy array and impute non-finite (inf/NaN) values. + # --------------------------------------------------------- + logger.info("Validating perturbation distances...") + distances_raw = np.array([d for _, d in raw_wmd_scores], dtype=float) + + # Filter out non-finite values to determine the maximum valid distance + valid_distances = distances_raw[np.isfinite(distances_raw)] + + # Calculate max_penalty: max valid distance + 1000, or default 1000 if + # all failed + if len(valid_distances) > 0: + max_penalty = np.max(valid_distances) + 1000.0 + else: + max_penalty = 1000.0 + + # Impute infinite/NaN values with the dynamically calculated maximum penalty + distances_array = np.where( + np.isfinite(distances_raw), distances_raw, max_penalty + ) + + # Reconstruct wmd_scores with imputed values for downstream consistency + wmd_scores = [ + (text, float(dist)) + for (text, _), dist in zip(raw_wmd_scores, distances_array, strict=False) + ] masks_as_arrays: list[np.ndarray] = [ np.array(m, dtype=int) for m in binary_masks diff --git a/src/xwhy/providers/anthropic.py b/src/xwhy/providers/anthropic.py index e8ff670a..e215da19 100644 --- a/src/xwhy/providers/anthropic.py +++ b/src/xwhy/providers/anthropic.py @@ -1,5 +1,8 @@ """Anthropic provider implementation.""" +import time +from typing import Any + from xwhy.logger import logger from xwhy.providers.base import BaseProvider @@ -24,58 +27,86 @@ def _generate( model: str, max_tokens: int, temperature: float, + **kwargs: Any, # noqa: ANN401 ) -> str: - """Generate text from Anthropic. + """Generate text from Anthropic with built-in retries. Args: prompt: Input prompt. model: Anthropic model name. max_tokens: Maximum output tokens. temperature: Sampling temperature. + **kwargs: Extra parameters (supports 'max_retries' and 'delay'). Returns: - Generated text. + Generated text string. Raises: - RuntimeError: If the API returns an empty response. + RuntimeError: If the API returns an empty response or fails + after all retries. """ - try: - response = self._client.messages.create( - model=model, - max_tokens=max_tokens, - temperature=temperature, - messages=[ - { - "role": "user", - "content": prompt, - } - ], - ) - - # Anthropic returns a list of ContentBlock objects. We extract the text - # from the first block if it exists to avoid IndexError. - result_text = "" - if response.content: - result_text = str(response.content[0].text).strip() - - if not result_text: - error_message = ( - "Received an empty response from the Anthropic API. " - "This could be due to content moderation filters, network" - " filtering (anti-filter), or provider-side anomalies." + max_retries: int = kwargs.get("max_retries", 7) + delay_override: float | None = kwargs.get("delay") + + for retry_number in range(1, max_retries + 1): + try: + response = self._client.messages.create( + model=model, + max_tokens=max_tokens, + temperature=temperature, + messages=[ + { + "role": "user", + "content": prompt, + } + ], ) - logger.error(error_message) - raise RuntimeError(error_message) - - return result_text - except RuntimeError: - raise + # Anthropic returns a list of ContentBlock objects. We extract + # the text from the first block if it exists to avoid IndexError. + result_text = "" + if response.content: + result_text = str(response.content[0].text).strip() + + if not result_text: + error_message = ( + "Received an empty response from the Anthropic API. " + "This could be due to content moderation filters, " + "network filtering (anti-filter), or " + "provider-side anomalies." + ) + logger.error(error_message) + raise RuntimeError(error_message) + + return result_text + + except RuntimeError: + raise + + except Exception as exc: + if retry_number == max_retries: + logger.error( + "Anthropic request failed after %d retries: %s", + max_retries, + exc, + ) + raise RuntimeError(f"Anthropic request failed: {exc}") from exc + + delay: float = ( + delay_override + if delay_override is not None + else min(2**retry_number, 30) + ) + logger.warning( + "Retry %d/%d for Anthropic text generation. Waiting %s seconds...", + retry_number, + max_retries, + delay, + ) + time.sleep(delay) - except Exception as exc: - logger.error("Anthropic request failed: %s", exc) - raise RuntimeError(f"Anthropic request failed: {exc}") from exc + raise RuntimeError("Anthropic text generation failed after max retries.") def answer( self, @@ -84,6 +115,7 @@ def answer( model: str = "claude-opus-4-8", max_tokens: int = 1024, temperature: float = 0.0, + **kwargs: Any, # noqa: ANN401 ) -> str: """Generate a natural-language answer. @@ -92,9 +124,10 @@ def answer( model: Anthropic model name. max_tokens: Maximum output tokens. temperature: Sampling temperature. + **kwargs: Extra parameters (supports 'max_retries' and 'delay'). Returns: - Generated response text. + Generated response text string. """ return self._generate( @@ -102,4 +135,5 @@ def answer( model=model, max_tokens=max_tokens, temperature=temperature, + **kwargs, ) diff --git a/src/xwhy/providers/gemini.py b/src/xwhy/providers/gemini.py index 2802fa5f..aeadf0de 100644 --- a/src/xwhy/providers/gemini.py +++ b/src/xwhy/providers/gemini.py @@ -35,62 +35,84 @@ def _generate( model: str, max_tokens: int, temperature: float, + **kwargs: Any, # noqa: ANN401 ) -> str: - """Generate text from Gemini. + """Generate text from Gemini with built-in retries. Args: prompt: Input prompt. model: Gemini model name. max_tokens: Maximum output tokens. temperature: Sampling temperature. + **kwargs: Extra parameters (supports 'max_retries' and 'delay'). Returns: Generated text. Raises: RuntimeError: If the API returns an empty response or is - blocked by safety filters. + blocked by safety filters after all retries. """ - try: - response = self._client.models.generate_content( - model=model, - contents=types.Part.from_text(text=prompt), - config=types.GenerateContentConfig( - max_output_tokens=max_tokens, - temperature=temperature, - ), - ) + max_retries: int = kwargs.get("max_retries", 7) + delay_override: float | None = kwargs.get("delay") + for retry_number in range(1, max_retries + 1): try: - result_text = str(response.text).strip() - except ValueError as val_err: - # Gemini throws ValueError on .text access if the response was - # blocked by safety filters. - error_message = ( - f"Gemini generation was blocked (likely due to safety filters) " - f"for model '{model}'. No content returned." - ) - logger.error(error_message) - raise RuntimeError(error_message) from val_err - - if not result_text: - error_message = ( - "Received an empty response from the Gemini API. " - "This could be due to network filtering (anti-filter) " - "or provider-side anomalies." + response = self._client.models.generate_content( + model=model, + contents=types.Part.from_text(text=prompt), + config=types.GenerateContentConfig( + max_output_tokens=max_tokens, + temperature=temperature, + ), ) - logger.error(error_message) - raise RuntimeError(error_message) - return result_text + try: + result_text = str(response.text).strip() + except ValueError as val_err: + # Gemini throws ValueError on .text access if the response was + # blocked by safety filters. + error_message = ( + f"Gemini generation was blocked (likely due to safety filters) " + f"for model '{model}'. No content returned." + ) + raise RuntimeError(error_message) from val_err + + if not result_text: + error_message = ( + "Received an empty response from the Gemini API. " + "This could be due to network filtering (anti-filter) " + "or provider-side anomalies." + ) + raise RuntimeError(error_message) + + return result_text + + except Exception as exc: + if retry_number == max_retries: + logger.error( + "Gemini request failed after %d retries: %s", max_retries, exc + ) + raise RuntimeError(f"Gemini request failed: {exc}") from exc - except RuntimeError: - raise + # Use provided delay or exponential backoff maxing at 30 seconds + delay: float = ( + delay_override + if delay_override is not None + else min(2**retry_number, 30) + ) + logger.warning( + "Retry %d/%d for Gemini text generation. Waiting %s seconds...", + retry_number, + max_retries, + delay, + ) + time.sleep(delay) - except Exception as exc: - logger.error("Gemini request failed: %s", exc) - raise RuntimeError(f"Gemini request failed: {exc}") from exc + # Catch-all to satisfy mypy in case max_retries <= 0 prevents the loop + # from running + raise RuntimeError("Failed to generate text: max_retries must be at least 1.") def answer( self, @@ -99,6 +121,7 @@ def answer( model: str = "gemini-2.5-flash", max_tokens: int = 200, temperature: float = 0.0, + **kwargs: Any, # noqa: ANN401 ) -> str: """Generate a natural-language answer. @@ -107,6 +130,7 @@ def answer( model: Gemini model name. max_tokens: Maximum output tokens. temperature: Sampling temperature. + **kwargs: Extra parameters (supports 'max_retries' and 'delay'). Returns: Generated response text. @@ -117,6 +141,7 @@ def answer( model=model, max_tokens=max_tokens, temperature=temperature, + **kwargs, ) # ------------------------------------------------------------------------- @@ -139,7 +164,7 @@ def _execute_image_request( input_image_path: str | None = None, **kwargs: Any, # noqa: ANN401 ) -> tuple[bool, str]: - """Execute the core generation logic for image requests. + """Execute the core generation logic for image requests with built-in retries. Args: prompt: The text prompt provided by the user. @@ -171,43 +196,71 @@ def _execute_image_request( ) generated_img: Image.Image | None = None - gen_img_flag = True final_mime = default_mime_type - try: - if stream: - response_iter = self._client.models.generate_content_stream( - model=model_name, - contents=contents, - config=generate_content_config, - ) - for chunk in response_iter: - for part in chunk.parts: + max_retries: int = kwargs.get("max_retries", 7) + delay_override: float | None = kwargs.get("delay") + + for retry_number in range(1, max_retries + 1): + try: + if stream: + response_iter = self._client.models.generate_content_stream( + model=model_name, + contents=contents, + config=generate_content_config, + ) + for chunk in response_iter: + for part in chunk.parts: + if part.inline_data is not None: + img_data = BytesIO(part.inline_data.data) + generated_img = Image.open(img_data) + final_mime = part.inline_data.mime_type + break + else: + response = self._client.models.generate_content( + model=model_name, + contents=contents, + config=generate_content_config, + ) + for part in response.parts: if part.inline_data is not None: img_data = BytesIO(part.inline_data.data) generated_img = Image.open(img_data) final_mime = part.inline_data.mime_type break - else: - response = self._client.models.generate_content( - model=model_name, - contents=contents, - config=generate_content_config, + except Exception as e: + logger.exception( + "Error during API call attempt %d: %s", retry_number, e + ) + + # If successful, exit the retry loop immediately + if generated_img is not None: + break + + # If failed and we have retries left, wait before trying again + if retry_number < max_retries: + # Use provided delay or exponential backoff maxing at 30 seconds + delay: float = ( + delay_override + if delay_override is not None + else min(2**retry_number, 30) + ) + logger.warning( + "Retry %s/%s for image generation. Waiting %s seconds...", + retry_number, + max_retries, + delay, ) - for part in response.parts: - if part.inline_data is not None: - img_data = BytesIO(part.inline_data.data) - generated_img = Image.open(img_data) - final_mime = part.inline_data.mime_type - break - except Exception as e: - logger.exception("Error during API call: %s", e) + time.sleep(delay) + # Retries exhausted. Check if we still don't have an image. if generated_img is None: gen_img_flag = False logger.debug( - "Failed to generate image for prompt: '%s'. Creating placeholder.", + "Failed to generate image for prompt: '%s' after %d retries. " + "Creating placeholder.", prompt, + max_retries, ) fallback_img = self._create_placeholder_image( prompt=prompt, output_dir=output_dir, save=False @@ -215,6 +268,8 @@ def _execute_image_request( if isinstance(fallback_img, Image.Image): generated_img = fallback_img final_mime = "image/png" + else: + gen_img_flag = True # Map supported Google Gemini MIME types to extensions mime_to_ext = { @@ -293,6 +348,7 @@ def generate_image( stream=stream, seed=seed, input_image_path=None, + **kwargs, ) def edit_image( @@ -360,6 +416,7 @@ def edit_image( seed=seed, default_mime_type=mime_type, input_image_path=image_path, + **kwargs, ) def submit_image_batch( diff --git a/src/xwhy/providers/huggingface.py b/src/xwhy/providers/huggingface.py index 3375b44a..2f5f601a 100644 --- a/src/xwhy/providers/huggingface.py +++ b/src/xwhy/providers/huggingface.py @@ -124,10 +124,13 @@ def _initialize_pipeline(self) -> Any: # noqa: ANN401 # General Case: Any HuggingFace Diffusers model logger.debug( - "Attempting to initialize Diffusers pipeline for '%s'...", self.model_name + "Attempting to initialize Diffusers pipeline for '%s'...", + self.model_name, ) try: - from diffusers.pipelines.auto_pipeline import AutoPipelineForText2Image + from diffusers.pipelines.auto_pipeline import ( + AutoPipelineForText2Image, + ) pipe = AutoPipelineForText2Image.from_pretrained( # type: ignore[no-untyped-call] self.model_name, @@ -140,13 +143,16 @@ def _initialize_pipeline(self) -> Any: # noqa: ANN401 except Exception as exc: logger.debug( - "Could not load '%s' as AutoPipelineForText2Image (This is expected" - " if it's an LLM). Fallback to DiffusionPipeline. Error: %s", + "Could not load '%s' as AutoPipelineForText2Image " + "(This is expected if it's an LLM). Fallback to " + "DiffusionPipeline. Error: %s", self.model_name, exc, ) try: - from diffusers.pipelines.pipeline_utils import DiffusionPipeline + from diffusers.pipelines.pipeline_utils import ( + DiffusionPipeline, + ) pipe = DiffusionPipeline.from_pretrained( # type: ignore[assignment] self.model_name, @@ -174,45 +180,73 @@ def _generate( model: str, max_tokens: int, temperature: float, + **kwargs: Any, # noqa: ANN401 ) -> str: - """Generate text from HuggingFace. + """Generate text from HuggingFace with built-in retries. Args: prompt: Input prompt. model: HuggingFace model name. max_tokens: Maximum output tokens. temperature: Sampling temperature. + **kwargs: Extra parameters (supports 'max_retries' and 'delay'). Returns: - Generated text. + Generated text string. Raises: - RuntimeError: If the API returns an empty response or fails. + RuntimeError: If the API returns an empty response or fails + after all retries. """ - try: - response = self._client.chat.completions.create( - model=model, - messages=[{"role": "user", "content": prompt}], - max_tokens=max_tokens, - temperature=temperature, - ) - - result_text = str(response.choices[0].message.content).strip() + max_retries: int = kwargs.get("max_retries", 7) + delay_override: float | None = kwargs.get("delay") - if not result_text: - error_message = ( - "Received an empty response from the HuggingFace API. " - "This could be due to guardrails or network filtering." + for retry_number in range(1, max_retries + 1): + try: + response = self._client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": prompt}], + max_tokens=max_tokens, + temperature=temperature, ) - logger.error(error_message) - raise RuntimeError(error_message) - return result_text + result_text = str(response.choices[0].message.content).strip() + + if not result_text: + error_message = ( + "Received an empty response from the HuggingFace API. " + "This could be due to guardrails or network filtering." + ) + logger.error(error_message) + raise RuntimeError(error_message) + + return result_text + + except Exception as exc: + if retry_number == max_retries: + logger.error( + "HuggingFace request failed after %d retries: %s", + max_retries, + exc, + ) + raise RuntimeError(f"HuggingFace request failed: {exc}") from exc + + delay: float = ( + delay_override + if delay_override is not None + else min(2**retry_number, 30) + ) + logger.warning( + "Retry %d/%d for HuggingFace text generation. Waiting %s " + "seconds...", + retry_number, + max_retries, + delay, + ) + time.sleep(delay) - except Exception as exc: - logger.error("HuggingFace request failed: %s", exc) - raise RuntimeError(f"HuggingFace request failed: {exc}") from exc + raise RuntimeError("HuggingFace text generation failed after max retries.") def answer( self, @@ -221,6 +255,7 @@ def answer( model: str = "meta-llama/Meta-Llama-3-8B-Instruct", max_tokens: int = 512, temperature: float = 0.1, + **kwargs: Any, # noqa: ANN401 ) -> str: """Generate a natural-language answer. @@ -229,9 +264,10 @@ def answer( model: HuggingFace model name. max_tokens: Maximum output tokens. temperature: Sampling temperature. + **kwargs: Extra parameters (supports 'max_retries' and 'delay'). Returns: - Generated response text. + Generated response text string. """ return self._generate( @@ -239,6 +275,7 @@ def answer( model=model, max_tokens=max_tokens, temperature=temperature, + **kwargs, ) # ------------------------------------------------------------------------- @@ -247,7 +284,7 @@ def answer( @property def supports_mask(self) -> bool: - """Check if the underlying diffusers pipeline supports/requires a mask image.""" + """Check if underlying diffusers pipeline supports/requires a mask image.""" if self.pipe is None: return False pipe_class_name = type(self.pipe).__name__ @@ -263,7 +300,7 @@ def _execute_image_request( class_names: Sequence[str] | None = None, **kwargs: Any, # noqa: ANN401 ) -> tuple[bool, str]: - """Execute core image generation or editing logic for HuggingFace pipelines. + """Execute core image generation or editing logic with built-in retries. Args: prompt: Text instruction describing the image modification. @@ -272,21 +309,22 @@ def _execute_image_request( segmentation_model: Optional model for generating segmentation masks. transform_fn: Preprocessing transform for PIL image. class_names: Optional sequence of class names for mask generation. - **kwargs: Additional pipeline-specific parameters. + **kwargs: Additional pipeline-specific parameters (supports + 'max_retries' and 'delay'). Returns: A tuple containing a boolean success flag and the generated file path. Raises: - RuntimeError: If the pipeline is not initialized. + RuntimeError: If the pipeline is not initialized or execution fails. ValueError: If an inpainting model is used without a segmentation model. """ if self.pipe is None: raise RuntimeError("HuggingFace pipeline is not initialized.") - gen_img_flag = True - generated_img: Image.Image | None = None + max_retries: int = kwargs.pop("max_retries", 7) + delay_override: float | None = kwargs.pop("delay", None) image: Image.Image | None = None if input_image_path is not None: @@ -313,47 +351,86 @@ def _execute_image_request( "segmentation_model is required for Inpainting pipelines." ) - try: - # Auto-adapt general Text2Image pipeline to Image2Image if editing - current_pipe = self.pipe - if ( - input_image_path is not None - and not is_pix2pix - and not is_inpaint - and "Image2Image" not in pipe_class_name - ): - from diffusers.pipelines.auto_pipeline import AutoPipelineForImage2Image - - logger.debug("Converting pipeline to AutoPipelineForImage2Image...") - current_pipe = AutoPipelineForImage2Image.from_pipe(self.pipe) # type: ignore[no-untyped-call] - - call_kwargs: dict[str, Any] = {"prompt": prompt} - if image is not None: - call_kwargs["image"] = image - if mask_image is not None and self.supports_mask: - call_kwargs["mask_image"] = mask_image - - if "num_inference_steps" not in kwargs: - kwargs["num_inference_steps"] = 50 if is_inpaint else 30 - - if is_pix2pix and "image_guidance_scale" not in kwargs: - kwargs["image_guidance_scale"] = 1.0 - - call_kwargs.update(kwargs) - - output = current_pipe(**call_kwargs) - - if hasattr(output, "images") and output.images: - generated_img = output.images[0] - elif isinstance(output, list) and output: - generated_img = output[0] - else: - raise RuntimeError("No valid images found in pipeline output.") + generated_img: Image.Image | None = None - except Exception as exc: + for retry_number in range(1, max_retries + 1): + try: + # Auto-adapt general Text2Image pipeline to Image2Image if editing + current_pipe = self.pipe + if ( + input_image_path is not None + and not is_pix2pix + and not is_inpaint + and "Image2Image" not in pipe_class_name + ): + from diffusers.pipelines.auto_pipeline import ( + AutoPipelineForImage2Image, + ) + + logger.debug("Converting pipeline to AutoPipelineForImage2Image...") + current_pipe = AutoPipelineForImage2Image.from_pipe(self.pipe) # type: ignore[no-untyped-call] + + call_kwargs: dict[str, Any] = {"prompt": prompt} + if image is not None: + call_kwargs["image"] = image + if mask_image is not None and self.supports_mask: + call_kwargs["mask_image"] = mask_image + + if "num_inference_steps" not in kwargs: + kwargs["num_inference_steps"] = 50 if is_inpaint else 30 + + if is_pix2pix and "image_guidance_scale" not in kwargs: + kwargs["image_guidance_scale"] = 1.0 + + call_kwargs.update(kwargs) + + output = current_pipe(**call_kwargs) + + if hasattr(output, "images") and output.images: + generated_img = output.images[0] + elif isinstance(output, list) and output: + generated_img = output[0] + else: + raise RuntimeError("No valid images found in pipeline output.") + + if generated_img is not None: + break + + except Exception as exc: + logger.warning( + "Pipeline execution failed on attempt %d/%d for prompt '%s': %s", + retry_number, + max_retries, + prompt, + exc, + ) + if retry_number == max_retries: + logger.exception( + "All HuggingFace image generation retries exhausted." + ) + + if generated_img is None and retry_number < max_retries: + delay: float = ( + delay_override + if delay_override is not None + else min(2**retry_number, 30) + ) + logger.warning( + "Retry %d/%d for image generation. Waiting %s seconds...", + retry_number, + max_retries, + delay, + ) + time.sleep(delay) + + gen_img_flag = True + if generated_img is None: gen_img_flag = False - logger.exception( - "Pipeline execution failed for prompt '%s': %s", prompt, exc + logger.debug( + "Failed to generate image for prompt: '%s' after %d retries. " + "Creating placeholder.", + prompt, + max_retries, ) fallback_img = self._create_placeholder_image( prompt=prompt, output_dir=output_dir, save=False @@ -387,7 +464,8 @@ def generate_image( Args: prompt: Text description of the desired image. output_dir: Directory where the image will be stored. - **kwargs: Extra parameters for the pipeline. + **kwargs: Extra parameters for the pipeline (supports 'max_retries' + and 'delay'). Returns: A tuple of a boolean success flag and the generated file path. @@ -420,7 +498,8 @@ def edit_image( segmentation_model: Optional model for generating segmentation masks. transform_fn: Preprocessing transform for PIL image. class_names: Optional sequence of class names for mask generation. - **kwargs: Extra parameters for the pipeline. + **kwargs: Extra parameters for the pipeline (supports 'max_retries' + and 'delay'). Returns: A tuple of a boolean success flag and the generated file path. diff --git a/src/xwhy/providers/openai.py b/src/xwhy/providers/openai.py index 4d986879..931542da 100644 --- a/src/xwhy/providers/openai.py +++ b/src/xwhy/providers/openai.py @@ -50,100 +50,133 @@ def _generate( model: str, max_tokens: int, temperature: float, + **kwargs: Any, # noqa: ANN401 ) -> str: - """Generate text from OpenAI. + """Generate text from OpenAI with built-in retries and error handling. Args: prompt: Input prompt. - model: OpenAI model. + model: OpenAI model identifier. max_tokens: Maximum output tokens. temperature: Sampling temperature. + **kwargs: Extra parameters (supports 'max_retries' and 'delay'). Returns: - Generated text. + Generated text string. Raises: - RuntimeError: If the API returns an empty response. + RuntimeError: If the API returns an empty response or fails + after all retries. """ - try: - if self._is_reasoning_model(model): - reasoning_response = self._client.responses.create( - model=model, - input=prompt, - max_output_tokens=max_tokens, - reasoning={"effort": "low"}, - temperature=temperature, - ) - result_text = str(reasoning_response.output_text).strip() - else: - completion_response = self._client.completions.create( - model=model, - prompt=prompt, - max_tokens=max_tokens, - temperature=temperature, - ) - result_text = str(completion_response.choices[0].text).strip() - - if not result_text: - error_message = ( - "Received an empty response from the OpenAI API. " - "This could be due to safety guardrails, network " - "filtering (anti-filter), or provider-side anomalies." - ) - logger.error(error_message) - raise RuntimeError(error_message) - - return result_text + max_retries: int = kwargs.get("max_retries", 7) + delay_override: float | None = kwargs.get("delay") - except RuntimeError: - raise - - except Exception as exc: - error_msg = str(exc).lower() - - if "temperature" in error_msg and ( - "support" in error_msg or "value" in error_msg or "allowed" in error_msg - ): - logger.warning( - "Dynamic fix applied: temperature=%f is not supported for model " - "'%s'. Retrying automatically with default temperature (1.0).", - temperature, - model, - ) - if temperature != 1.0: - return self._generate( - prompt=prompt, + for retry_number in range(1, max_retries + 1): + try: + if self._is_reasoning_model(model): + reasoning_response = self._client.responses.create( + model=model, + input=prompt, + max_output_tokens=max_tokens, + reasoning={"effort": "low"}, + temperature=temperature, + ) + result_text = str(reasoning_response.output_text).strip() + else: + completion_response = self._client.completions.create( model=model, + prompt=prompt, max_tokens=max_tokens, - temperature=1.0, + temperature=temperature, + ) + result_text = str(completion_response.choices[0].text).strip() + + if not result_text: + error_message = ( + "Received an empty response from the OpenAI API. " + "This could be due to safety guardrails, network " + "filtering (anti-filter), or provider-side anomalies." ) + logger.error(error_message) + raise RuntimeError(error_message) + + return result_text + + except RuntimeError: + raise - if ( - "max_output_tokens" in error_msg - and "integer below minimum value" in error_msg - ): - match = re.search(r"expected a value >= (\d+)", error_msg) + except Exception as exc: + error_msg = str(exc).lower() - if match: - required_min = int(match.group(1)) + if "temperature" in error_msg and ( + "support" in error_msg + or "value" in error_msg + or "allowed" in error_msg + ): logger.warning( - "Dynamic fix applied: max_tokens=%d is too low for model '%s'. " - "Retrying automatically with required minimum: %d.", - max_tokens, + "Dynamic fix applied: temperature=%f is not supported " + "for model '%s'. Retrying automatically with default " + "temperature (1.0).", + temperature, model, - required_min, ) - - return self._generate( - prompt=prompt, - model=model, - max_tokens=required_min, - temperature=temperature, + if temperature != 1.0: + return self._generate( + prompt=prompt, + model=model, + max_tokens=max_tokens, + temperature=1.0, + **kwargs, + ) + + if ( + "max_output_tokens" in error_msg + and "integer below minimum value" in error_msg + ): + match = re.search(r"expected a value >= (\d+)", error_msg) + + if match: + required_min = int(match.group(1)) + logger.warning( + "Dynamic fix applied: max_tokens=%d is too low " + "for model '%s'. Retrying automatically with " + "required minimum: %d.", + max_tokens, + model, + required_min, + ) + + return self._generate( + prompt=prompt, + model=model, + max_tokens=required_min, + temperature=temperature, + **kwargs, + ) + + if retry_number == max_retries: + logger.error( + "OpenAI request failed after %d retries: %s", + max_retries, + exc, ) + raise RuntimeError(f"OpenAI request failed: {exc}") from exc + + delay: float = ( + delay_override + if delay_override is not None + else min(2**retry_number, 30) + ) + logger.warning( + "Retry %d/%d for OpenAI text generation. Waiting %s seconds...", + retry_number, + max_retries, + delay, + ) + time.sleep(delay) - logger.error("OpenAI request failed: %s", exc) - raise RuntimeError(f"OpenAI request failed: {exc}") from exc + raise RuntimeError("OpenAI text generation failed after max retries.") def answer( self, @@ -152,6 +185,7 @@ def answer( model: str = "gpt-3.5-turbo-instruct", max_tokens: int = 200, temperature: float = 0.0, + **kwargs: Any, # noqa: ANN401 ) -> str: """Generate a natural-language answer. @@ -160,9 +194,10 @@ def answer( model: OpenAI model name. max_tokens: Maximum output tokens. temperature: Sampling temperature. + **kwargs: Extra parameters (supports 'max_retries' and 'delay'). Returns: - Generated response text. + Generated response text string. """ return self._generate( @@ -170,6 +205,7 @@ def answer( model=model, max_tokens=max_tokens, temperature=temperature, + **kwargs, ) # ------------------------------------------------------------------------- @@ -191,18 +227,25 @@ def _execute_image_request( output_dir: Directory to save the final generated output. model_name: Target OpenAI model identifier. input_image_path: Path to base image if editing, None otherwise. - **kwargs: Additional parameters (e.g., extra_body, response_format). + **kwargs: Additional parameters (e.g., extra_body, response_format, + max_retries, delay). Returns: A tuple containing a boolean success flag and the file path. + Raises: + RuntimeError: If image data is missing or empty. + FileNotFoundError: If the input image path does not exist. + """ - # Safely extract configuration flags to prevent them from reaching the API - provider_name = kwargs.pop("provider_name", "openai") - output_format = kwargs.pop("output_format", "png") - use_generate_for_edit = kwargs.pop("use_generate_for_edit", False) - use_image_data_uri = kwargs.pop("use_image_data_uri", False) - use_image_url = kwargs.pop("use_image_url", False) + # Safely extract configuration flags to prevent them from reaching API + provider_name: str = kwargs.pop("provider_name", "openai") + output_format: str = kwargs.pop("output_format", "png") + use_generate_for_edit: bool = kwargs.pop("use_generate_for_edit", False) + use_image_data_uri: bool = kwargs.pop("use_image_data_uri", False) + use_image_url: bool = kwargs.pop("use_image_url", False) + max_retries: int = kwargs.pop("max_retries", 7) + delay_override: float | None = kwargs.pop("delay", None) # Handle response_format logic cleanly if "response_format" in kwargs and kwargs["response_format"] is None: @@ -211,70 +254,97 @@ def _execute_image_request( kwargs.setdefault("response_format", "b64_json") generated_img: Image.Image | None = None - gen_img_flag = True - try: - # Standard OpenAI Edit Request - if input_image_path is not None and not use_generate_for_edit: - with open(input_image_path, "rb") as img_file: - response = self._client.images.edit( + for retry_number in range(1, max_retries + 1): + try: + # Standard OpenAI Edit Request + if input_image_path is not None and not use_generate_for_edit: + with open(input_image_path, "rb") as img_file: + response = self._client.images.edit( + model=model_name, + image=img_file, + prompt=prompt, + **kwargs, + ) + else: + # Handle ByteDance-style image-to-image via generate + if input_image_path is not None and use_generate_for_edit: + input_image_data_uri = image_to_base64( + image_path=input_image_path, + include_data_uri=use_image_data_uri, + ) + + extra_body = kwargs.pop("extra_body", {}) + + if use_image_url: + if "image_url" not in extra_body: + extra_body["image_url"] = input_image_data_uri + else: + if "image" not in extra_body: + extra_body["image"] = input_image_data_uri + + kwargs["extra_body"] = extra_body + + # Standard or Alternative Generation Request + response = self._client.images.generate( model=model_name, - image=img_file, prompt=prompt, **kwargs, ) - else: - # Handle ByteDance-style image-to-image via generate endpoint - if input_image_path is not None and use_generate_for_edit: - input_image_data_uri = image_to_base64( - image_path=input_image_path, - include_data_uri=use_image_data_uri, - ) - extra_body = kwargs.pop("extra_body", {}) - - if use_image_url: - if "image_url" not in extra_body: - extra_body["image_url"] = input_image_data_uri + # Support both b64_json and url formats seamlessly + if response.data: + img_data_obj = response.data[0] + if hasattr(img_data_obj, "b64_json") and img_data_obj.b64_json: + img_bytes = BytesIO(base64.b64decode(img_data_obj.b64_json)) + generated_img = Image.open(img_bytes) + elif hasattr(img_data_obj, "url") and img_data_obj.url: + req_response = requests.get(img_data_obj.url, timeout=30) + req_response.raise_for_status() + generated_img = Image.open(BytesIO(req_response.content)) else: - if "image" not in extra_body: - extra_body["image"] = input_image_data_uri + raise RuntimeError( + "No valid image data (b64_json or url) found." + ) + else: + raise RuntimeError("Empty image data returned from provider.") - kwargs["extra_body"] = extra_body + if generated_img is not None: + break - # Standard or Alternative Generation Request - response = self._client.images.generate( - model=model_name, - prompt=prompt, - **kwargs, + except Exception as e: + logger.warning( + "Error during OpenAI API call on attempt %d/%d: %s", + retry_number, + max_retries, + e, ) + if retry_number == max_retries: + logger.exception("All image generation retries exhausted.") + + if generated_img is None and retry_number < max_retries: + delay = ( + delay_override + if delay_override is not None + else min(2**retry_number, 30) + ) + logger.warning( + "Retry %d/%d for image generation. Waiting %s seconds...", + retry_number, + max_retries, + delay, + ) + time.sleep(delay) - # Support both b64_json and url formats seamlessly - if response.data: - img_data_obj = response.data[0] - if hasattr(img_data_obj, "b64_json") and img_data_obj.b64_json: - img_bytes = BytesIO(base64.b64decode(img_data_obj.b64_json)) - generated_img = Image.open(img_bytes) - elif hasattr(img_data_obj, "url") and img_data_obj.url: - req_response = requests.get(img_data_obj.url, timeout=30) - req_response.raise_for_status() - generated_img = Image.open(BytesIO(req_response.content)) - else: - raise RuntimeError( - "No valid image data (b64_json or url) found in response." - ) - else: - raise RuntimeError("Empty image data returned from provider.") - - except Exception as e: - logger.exception("Error during OpenAI-compatible API call: %s", e) - - # Fallback to placeholder if everything failed + gen_img_flag = True + # Fallback to placeholder if everything failed after all retries if generated_img is None: gen_img_flag = False logger.debug( - "Failed to generate image for prompt: '%s'. Creating placeholder.", + "Failed to generate image for prompt: '%s' after %d retries. " + "Creating placeholder.", prompt, + max_retries, ) if hasattr(self, "_create_placeholder_image"): fallback_img = self._create_placeholder_image( @@ -318,7 +388,8 @@ def generate_image( prompt: Text description of the desired image. output_dir: Directory where the image will be stored. model_name: OpenAI model name for image generation. - **kwargs: Extra parameters (e.g., output_format, extra_body). + **kwargs: Extra parameters (e.g., output_format, extra_body, + max_retries, delay). Returns: A tuple of a boolean success flag and the generated file path. @@ -347,11 +418,8 @@ def edit_image( image_path: Path to the source image file. output_dir: Directory where the edited image will be saved. model_name: OpenAI model name for image editing. - size: Sampling size specification for the output. - quality: Quality configuration ("low", "medium", "high", "auto"). - n: Number of output images to generate. **kwargs: Extra parameters (e.g., output_format, extra_body, - use_generate_for_edit). + use_generate_for_edit, max_retries, delay). Returns: A tuple of a boolean success flag and the generated file path. diff --git a/tests/explainers/test_image.py b/tests/explainers/test_image.py index 08d3cc1d..862ead9f 100644 --- a/tests/explainers/test_image.py +++ b/tests/explainers/test_image.py @@ -1037,6 +1037,156 @@ def classification_model(self) -> MagicMock | None: ImageClassificationExplainer.explain(explainer, "test.jpg") +@patch("xwhy.explainers.image.SurrogateTrainer") +@patch("xwhy.explainers.image.SurrogateFactory") +@patch("xwhy.explainers.image.RegressionMetrics") +@patch("xwhy.explainers.image.load_image_as_tensor") +@patch("xwhy.explainers.image.tensor_to_numpy_image") +def test_image_classification_explain_impute_when_some_distances_valid( + mock_tensor_to_np: MagicMock, + mock_load: MagicMock, + mock_metrics: MagicMock, + mock_factory: MagicMock, + mock_trainer: MagicMock, +) -> None: + """Cover the branch where at least one distance is finite. + + ``max_penalty`` becomes ``max(valid) + 1000.0`` and non-finite values + are replaced by that penalty. + """ + # Minimal explainer with mocked state + explainer = MagicMock(spec=ImageClassificationExplainer) + explainer.config = MagicMock() + explainer.config.use_best_surrogate = False + explainer.config.surrogate_type = SurrogateType.LIME + explainer.config.seed = 42 + explainer.config.num_top_features = 2 + explainer.config.num_top_predictions = 1 + explainer.config.use_model_preprocess = False + explainer.config.device = "cpu" + explainer.config.num_perturb = 3 + + explainer.state = MagicMock() + explainer.state.classification_model = MagicMock() + explainer.state.classification_model.preprocess_fn.mean = [0.0] + explainer.state.classification_model.preprocess_fn.std = [1.0] + explainer.state.classification_model.weights.meta = {"categories": ["cat", "dog"]} + explainer.state.classification_model.predict.return_value = torch.tensor( + [[0.1, 0.9]] + ) + explainer.state.transform_fn = None + explainer.state.perturbator = MagicMock() + explainer.state.perturbator.generate_superpixels.return_value = ( + np.zeros((8, 8), dtype=int), + 4, + ) + explainer.state.perturbator.generate.return_value = np.array( + [[1, 0, 1, 0], [0, 1, 0, 1], [1, 1, 0, 0]] + ) + explainer.state.perturbator.apply_mask.return_value = np.zeros((8, 8, 3)) + explainer.state.segmentation_model = None + explainer.state.embedding_model = None + + mock_load.return_value = (torch.zeros(1, 3, 8, 8), MagicMock()) + mock_tensor_to_np.return_value = np.zeros((8, 8, 3)) + + # Two finite + one non-finite distance + explainer._run_perturbation_loop = MagicMock( + return_value=( + np.array([[0.2, 0.8], [0.3, 0.7], [0.4, 0.6]]), + np.array([0.5, np.inf, 1.5]), + ) + ) + + mock_trainer.compute_weights.return_value = np.ones(3) + mock_surrogate = MagicMock() + mock_surrogate.coefficients.return_value = np.array([0.1, 0.2, 0.3, 0.4]) + mock_surrogate.predict.return_value = np.array([0.5, 0.6, 0.7]) + mock_factory.create.return_value = mock_surrogate + mock_metrics.calculate.return_value = MagicMock() + + # Bind the real explain method + result = ImageClassificationExplainer.explain(explainer, instance="dummy.png") + + distances = result.raw_data["distances"] + assert distances[0] == pytest.approx(0.5) + assert distances[1] == pytest.approx(1001.5) # max(0.5,1.5)+1000 + assert distances[2] == pytest.approx(1.5) + + +@patch("xwhy.explainers.image.SurrogateTrainer") +@patch("xwhy.explainers.image.SurrogateFactory") +@patch("xwhy.explainers.image.RegressionMetrics") +@patch("xwhy.explainers.image.load_image_as_tensor") +@patch("xwhy.explainers.image.tensor_to_numpy_image") +def test_image_classification_explain_impute_when_all_distances_non_finite( + mock_tensor_to_np: MagicMock, + mock_load: MagicMock, + mock_metrics: MagicMock, + mock_factory: MagicMock, + mock_trainer: MagicMock, +) -> None: + """Cover the branch where every distance is non-finite. + + ``max_penalty`` falls back to the constant ``1000.0``. + """ + explainer = MagicMock(spec=ImageClassificationExplainer) + explainer.config = MagicMock() + explainer.config.use_best_surrogate = False + explainer.config.surrogate_type = SurrogateType.LIME + explainer.config.seed = 42 + explainer.config.num_top_features = 2 + explainer.config.num_top_predictions = 1 + explainer.config.use_model_preprocess = False + explainer.config.device = "cpu" + explainer.config.num_perturb = 2 + + explainer.state = MagicMock() + explainer.state.classification_model = MagicMock() + explainer.state.classification_model.preprocess_fn.mean = [0.0] + explainer.state.classification_model.preprocess_fn.std = [1.0] + explainer.state.classification_model.weights.meta = {"categories": ["cat", "dog"]} + explainer.state.classification_model.predict.return_value = torch.tensor( + [[0.1, 0.9]] + ) + explainer.state.transform_fn = None + explainer.state.perturbator = MagicMock() + explainer.state.perturbator.generate_superpixels.return_value = ( + np.zeros((8, 8), dtype=int), + 4, + ) + explainer.state.perturbator.generate.return_value = np.array( + [[1, 0, 1, 0], [0, 1, 0, 1]] + ) + explainer.state.perturbator.apply_mask.return_value = np.zeros((8, 8, 3)) + explainer.state.segmentation_model = None + explainer.state.embedding_model = None + + mock_load.return_value = (torch.zeros(1, 3, 8, 8), MagicMock()) + mock_tensor_to_np.return_value = np.zeros((8, 8, 3)) + + # All non-finite + explainer._run_perturbation_loop = MagicMock( + return_value=( + np.array([[0.2, 0.8], [0.3, 0.7]]), + np.array([np.inf, np.nan]), + ) + ) + + mock_trainer.compute_weights.return_value = np.ones(2) + mock_surrogate = MagicMock() + mock_surrogate.coefficients.return_value = np.array([0.1, 0.2, 0.3, 0.4]) + mock_surrogate.predict.return_value = np.array([0.5, 0.6]) + mock_factory.create.return_value = mock_surrogate + mock_metrics.calculate.return_value = MagicMock() + + result = ImageClassificationExplainer.explain(explainer, instance="dummy.png") + + distances = result.raw_data["distances"] + assert distances[0] == pytest.approx(1000.0) + assert distances[1] == pytest.approx(1000.0) + + # ------------------------------------------------------------------------- # Image generation and editing # ------------------------------------------------------------------------- @@ -2294,3 +2444,148 @@ def test_generate_images_no_segmentation_model( ) assert len(res) == 2 assert all(success for success, _ in res) + + +@patch("xwhy.explainers.image.SurrogateTrainer") +@patch("xwhy.explainers.image.SurrogateFactory") +@patch("xwhy.explainers.image.RegressionMetrics") +@patch("xwhy.explainers.image.WMDDistance") +@patch("xwhy.explainers.image.DistanceNormalizer") +@patch("xwhy.explainers.image.save_data_to_pickle") +@patch("xwhy.explainers.image.save_perturbation_data_to_csv") +def test_image_generation_explain_impute_when_some_distances_valid( + mock_csv: MagicMock, + mock_pickle: MagicMock, + mock_normalizer: MagicMock, + mock_wmd: MagicMock, + mock_metrics: MagicMock, + mock_factory: MagicMock, + mock_trainer: MagicMock, +) -> None: + """Cover the branch where at least one image distance is finite.""" + explainer = MagicMock(spec=ImageGenerationAndEditingExplainer) + explainer.config = MagicMock() + explainer.config.use_best_surrogate = False + explainer.config.surrogate_type = SurrogateType.LIME + explainer.config.seed = 42 + explainer.config.num_perturbations = 3 + explainer.config.model_name = "test-model" + explainer.config.output_dir = "/tmp" + explainer.config.distance_type = MagicMock() + explainer.config.kernel_width = 0.25 + explainer.config.ridge_alpha = 1.0 + + explainer.state = MagicMock() + explainer.state.text_perturbator = MagicMock() + explainer.state.text_perturbator.generate.return_value = ( + ["p1", "p2", "p3"], + [[1, 0], [0, 1], [1, 1]], + ) + explainer.state.text_embedding_model = MagicMock() + explainer.state.engine = MagicMock() + + explainer._prepare_environment = MagicMock() + explainer._generate_images = MagicMock( + side_effect=[ + [(True, "/tmp/base.png")], # base + [(True, "/tmp/1.png"), (True, "/tmp/2.png"), (True, "/tmp/3.png")], + ] + ) + # Two finite + one inf + explainer._compute_perturbation_distances = MagicMock( + return_value=np.array([0.5, np.inf, 1.5]) + ) + + mock_wmd.return_value.compute_batch.return_value = [ + ("p1", 0.1), + ("p2", 0.2), + ("p3", 0.3), + ] + mock_normalizer.min_max.return_value = [("v", 0.5)] * 3 + mock_trainer.compute_weights.return_value = np.ones(3) + mock_surrogate = MagicMock() + mock_surrogate.coefficients.return_value = np.array([0.1, 0.2]) + mock_surrogate.predict.return_value = np.array([0.5, 0.6, 0.7]) + mock_factory.create.return_value = mock_surrogate + mock_metrics.calculate.return_value = MagicMock() + mock_csv.return_value = "/tmp/out.csv" + + result = ImageGenerationAndEditingExplainer.explain( + explainer, instance="a short descriptive prompt here" + ) + + y_target = result.raw_data["y_target"] + assert y_target[0] == pytest.approx(0.5) + assert y_target[1] == pytest.approx(1001.5) + assert y_target[2] == pytest.approx(1.5) + + +@patch("xwhy.explainers.image.SurrogateTrainer") +@patch("xwhy.explainers.image.SurrogateFactory") +@patch("xwhy.explainers.image.RegressionMetrics") +@patch("xwhy.explainers.image.WMDDistance") +@patch("xwhy.explainers.image.DistanceNormalizer") +@patch("xwhy.explainers.image.save_data_to_pickle") +@patch("xwhy.explainers.image.save_perturbation_data_to_csv") +def test_image_generation_explain_impute_when_all_distances_non_finite( + mock_csv: MagicMock, + mock_pickle: MagicMock, + mock_normalizer: MagicMock, + mock_wmd: MagicMock, + mock_metrics: MagicMock, + mock_factory: MagicMock, + mock_trainer: MagicMock, +) -> None: + """Cover the branch where every image distance is non-finite.""" + explainer = MagicMock(spec=ImageGenerationAndEditingExplainer) + explainer.config = MagicMock() + explainer.config.use_best_surrogate = False + explainer.config.surrogate_type = SurrogateType.LIME + explainer.config.seed = 42 + explainer.config.num_perturbations = 2 + explainer.config.model_name = "test-model" + explainer.config.output_dir = "/tmp" + explainer.config.distance_type = MagicMock() + explainer.config.kernel_width = 0.25 + explainer.config.ridge_alpha = 1.0 + + explainer.state = MagicMock() + explainer.state.text_perturbator = MagicMock() + explainer.state.text_perturbator.generate.return_value = ( + ["p1", "p2"], + [[1, 0], [0, 1]], + ) + explainer.state.text_embedding_model = MagicMock() + explainer.state.engine = MagicMock() + + explainer._prepare_environment = MagicMock() + explainer._generate_images = MagicMock( + side_effect=[ + [(True, "/tmp/base.png")], + [(True, "/tmp/1.png"), (True, "/tmp/2.png")], + ] + ) + explainer._compute_perturbation_distances = MagicMock( + return_value=np.array([np.inf, np.inf]) + ) + + mock_wmd.return_value.compute_batch.return_value = [ + ("p1", 0.1), + ("p2", 0.2), + ] + mock_normalizer.min_max.return_value = [("v", 0.5)] * 2 + mock_trainer.compute_weights.return_value = np.ones(2) + mock_surrogate = MagicMock() + mock_surrogate.coefficients.return_value = np.array([0.1]) + mock_surrogate.predict.return_value = np.array([0.5, 0.5]) + mock_factory.create.return_value = mock_surrogate + mock_metrics.calculate.return_value = MagicMock() + mock_csv.return_value = "/tmp/out.csv" + + result = ImageGenerationAndEditingExplainer.explain( + explainer, instance="a short descriptive prompt here" + ) + + y_target = result.raw_data["y_target"] + assert y_target[0] == pytest.approx(1000.0) + assert y_target[1] == pytest.approx(1000.0) diff --git a/tests/explainers/test_llm.py b/tests/explainers/test_llm.py index d5782bfd..93e4c8d7 100644 --- a/tests/explainers/test_llm.py +++ b/tests/explainers/test_llm.py @@ -328,3 +328,144 @@ def test_explain_fidelity_plot_flag( mock_plot.reset_mock() explainer.explain("test prompt") mock_plot.assert_not_called() + + +@patch("xwhy.explainers.llm.ProviderResolver.resolve") +@patch("xwhy.explainers.llm.TextPerturbation") +@patch("xwhy.explainers.llm.EmbeddingFactory") +@patch("xwhy.explainers.llm.WMDDistance") +@patch("xwhy.explainers.llm.DistanceNormalizer") +@patch("xwhy.explainers.llm.SurrogateTrainer") +@patch("xwhy.explainers.llm.SurrogateFactory") +@patch("xwhy.explainers.llm.RegressionMetrics") +def test_llm_explain_impute_when_some_distances_valid( + mock_metrics: MagicMock, + mock_surrogate_factory: MagicMock, + mock_trainer: MagicMock, + mock_normalizer: MagicMock, + mock_wmd: MagicMock, + mock_embedding_factory: MagicMock, + mock_perturbation: MagicMock, + mock_resolve: MagicMock, + mock_provider: MagicMock, +) -> None: + """Cover the branch where at least one WMD distance is finite. + + ``max_penalty`` must become ``max(valid) + 1000.0`` and every non-finite + value is replaced by that penalty. + + Args: + mock_metrics: Mock for RegressionMetrics. + mock_surrogate_factory: Mock for SurrogateFactory. + mock_trainer: Mock for SurrogateTrainer. + mock_normalizer: Mock for DistanceNormalizer. + mock_wmd: Mock for WMDDistance. + mock_embedding_factory: Mock for EmbeddingFactory. + mock_perturbation: Mock for TextPerturbation. + mock_resolve: Mock for ProviderResolver.resolve. + mock_provider: Fixture providing a mock BaseProvider. + + """ + mock_resolve.return_value = mock_provider + explainer = LLMExplainer(provider="openai", use_best_surrogate=False) + + mock_perturbation.return_value.generate.return_value = ( + ["res1", "res2", "res3"], + [np.array([1, 0]), np.array([0, 1]), np.array([1, 1])], + ) + mock_embedding_factory.create.return_value.load.return_value = MagicMock() + + # Two finite distances + one non-finite → valid branch is taken. + mock_wmd.return_value.compute_batch.return_value = [ + ("res1", 0.5), + ("res2", np.inf), + ("res3", 1.5), + ] + mock_normalizer.min_max.return_value = [ + ("val", 0.5), + ("val", 0.0), + ("val", 1.0), + ] + + mock_trainer.compute_weights.return_value = np.array([1.0, 1.0, 1.0]) + mock_surrogate = MagicMock() + mock_surrogate.coefficients.return_value = np.array([0.1, 0.2]) + mock_surrogate.predict.return_value = np.array([0.5, 0.6, 0.7]) + mock_surrogate_factory.create.return_value = mock_surrogate + + result = explainer.explain("test prompt") + + # max(0.5, 1.5) + 1000 = 1001.5 must have been used for the inf entry + distances_used = result.raw_data["wmd_scores"] + assert distances_used[0][1] == pytest.approx(0.5) + assert distances_used[1][1] == pytest.approx(1001.5) + assert distances_used[2][1] == pytest.approx(1.5) + assert isinstance(result, TextXWhyResult) + + +@patch("xwhy.explainers.llm.ProviderResolver.resolve") +@patch("xwhy.explainers.llm.TextPerturbation") +@patch("xwhy.explainers.llm.EmbeddingFactory") +@patch("xwhy.explainers.llm.WMDDistance") +@patch("xwhy.explainers.llm.DistanceNormalizer") +@patch("xwhy.explainers.llm.SurrogateTrainer") +@patch("xwhy.explainers.llm.SurrogateFactory") +@patch("xwhy.explainers.llm.RegressionMetrics") +def test_llm_explain_impute_when_all_distances_non_finite( + mock_metrics: MagicMock, + mock_surrogate_factory: MagicMock, + mock_trainer: MagicMock, + mock_normalizer: MagicMock, + mock_wmd: MagicMock, + mock_embedding_factory: MagicMock, + mock_perturbation: MagicMock, + mock_resolve: MagicMock, + mock_provider: MagicMock, +) -> None: + """Cover the branch where every WMD distance is non-finite. + + ``max_penalty`` must fall back to the constant ``1000.0``. + + Args: + mock_metrics: Mock for RegressionMetrics. + mock_surrogate_factory: Mock for SurrogateFactory. + mock_trainer: Mock for SurrogateTrainer. + mock_normalizer: Mock for DistanceNormalizer. + mock_wmd: Mock for WMDDistance. + mock_embedding_factory: Mock for EmbeddingFactory. + mock_perturbation: Mock for TextPerturbation. + mock_resolve: Mock for ProviderResolver.resolve. + mock_provider: Fixture providing a mock BaseProvider. + + """ + mock_resolve.return_value = mock_provider + explainer = LLMExplainer(provider="openai", use_best_surrogate=False) + + mock_perturbation.return_value.generate.return_value = ( + ["res1", "res2"], + [np.array([1, 0]), np.array([0, 1])], + ) + mock_embedding_factory.create.return_value.load.return_value = MagicMock() + + # All non-finite → else branch (max_penalty = 1000.0) + mock_wmd.return_value.compute_batch.return_value = [ + ("res1", np.inf), + ("res2", np.nan), + ] + mock_normalizer.min_max.return_value = [ + ("val", 0.0), + ("val", 0.0), + ] + + mock_trainer.compute_weights.return_value = np.array([1.0, 1.0]) + mock_surrogate = MagicMock() + mock_surrogate.coefficients.return_value = np.array([0.1]) + mock_surrogate.predict.return_value = np.array([0.5, 0.5]) + mock_surrogate_factory.create.return_value = mock_surrogate + + result = explainer.explain("test prompt") + + distances_used = result.raw_data["wmd_scores"] + assert distances_used[0][1] == pytest.approx(1000.0) + assert distances_used[1][1] == pytest.approx(1000.0) + assert isinstance(result, TextXWhyResult) diff --git a/tests/explainers/test_tabular.py b/tests/explainers/test_tabular.py index bce3eed6..9fdde656 100644 --- a/tests/explainers/test_tabular.py +++ b/tests/explainers/test_tabular.py @@ -230,3 +230,85 @@ def test_tabular_explainer_run_valid_delegation( # Assert the return value matches what explain returned assert result == mock_explain.return_value + + +@patch("xwhy.explainers.tabular.SurrogateTrainer") +@patch("xwhy.explainers.tabular.SurrogateFactory") +@patch("xwhy.explainers.tabular.RegressionMetrics") +@patch("xwhy.explainers.tabular.calculate_distance") +def test_tabular_explain_impute_when_some_distances_valid( + mock_calc_dist: MagicMock, + mock_metrics: MagicMock, + mock_factory: MagicMock, + mock_trainer: MagicMock, +) -> None: + """Cover the branch where at least one scaled distance is finite.""" + model = MagicMock() + model.predict.return_value = np.array([0, 1, 0]) + + explainer = TabularExplainer( + model=model, + num_perturbations=3, + num_distribution_samples=5, + use_best_surrogate=False, + seed=42, + validate_normalization=False, + ) + + # Force some non-finite distances inside the loop + # by making calculate_distance return mixed values + mock_calc_dist.side_effect = [0.5, 0.5, np.inf, 0.5, 1.5, 1.5] * 10 + + mock_trainer.compute_weights.return_value = np.ones(3) + mock_surrogate = MagicMock() + mock_surrogate.coefficients.return_value = np.array([0.1, 0.2]) + mock_surrogate.predict.return_value = np.array([0.5, 0.6, 0.7]) + mock_factory.create.return_value = mock_surrogate + mock_metrics.calculate.return_value = MagicMock() + + instance = np.array([0.1, -0.2]) + result = explainer.explain(instance) + + distances = result.raw_data["distances"] + # At least one entry must have been imputed with max+1000 + assert np.any(distances > 1000.0) + assert np.all(np.isfinite(distances)) + + +@patch("xwhy.explainers.tabular.SurrogateTrainer") +@patch("xwhy.explainers.tabular.SurrogateFactory") +@patch("xwhy.explainers.tabular.RegressionMetrics") +@patch("xwhy.explainers.tabular.calculate_distance") +def test_tabular_explain_impute_when_all_distances_non_finite( + mock_calc_dist: MagicMock, + mock_metrics: MagicMock, + mock_factory: MagicMock, + mock_trainer: MagicMock, +) -> None: + """Cover the branch where every scaled distance is non-finite.""" + model = MagicMock() + model.predict.return_value = np.array([0, 1, 0]) + + explainer = TabularExplainer( + model=model, + num_perturbations=2, + num_distribution_samples=3, + use_best_surrogate=False, + seed=42, + validate_normalization=False, + ) + + mock_calc_dist.return_value = np.inf + + mock_trainer.compute_weights.return_value = np.ones(2) + mock_surrogate = MagicMock() + mock_surrogate.coefficients.return_value = np.array([0.1, 0.2]) + mock_surrogate.predict.return_value = np.array([0.5, 0.5]) + mock_factory.create.return_value = mock_surrogate + mock_metrics.calculate.return_value = MagicMock() + + instance = np.array([0.0, 0.0]) + result = explainer.explain(instance) + + distances = result.raw_data["distances"] + assert np.allclose(distances, 1000.0) diff --git a/tests/explainers/test_text.py b/tests/explainers/test_text.py index 70a881bb..3243f8ce 100644 --- a/tests/explainers/test_text.py +++ b/tests/explainers/test_text.py @@ -514,3 +514,87 @@ def mock_predict_empty(texts: Sequence[str]) -> np.ndarray: ) assert isinstance(result, TextXWhyResult) + + +@patch("xwhy.explainers.text.SurrogateTrainer") +@patch("xwhy.explainers.text.SurrogateFactory") +@patch("xwhy.explainers.text.RegressionMetrics") +@patch("xwhy.explainers.text.WMDDistance") +def test_text_explain_impute_when_some_distances_valid( + mock_wmd: MagicMock, + mock_metrics: MagicMock, + mock_factory: MagicMock, + mock_trainer: MagicMock, +) -> None: + """Cover the branch where at least one WMD distance is finite.""" + predict_fn = MagicMock(return_value=np.array([[0.2, 0.8], [0.3, 0.7], [0.4, 0.6]])) + + with ( + patch("xwhy.explainers.text.EmbeddingFactory"), + patch("xwhy.explainers.text.TextPerturbation") as mock_pert, + ): + mock_pert.return_value.generate.return_value = ( + ["t1", "t2", "t3"], + [[1, 0], [0, 1], [1, 1]], + ) + explainer = TextExplainer(predict_fn=predict_fn, use_best_surrogate=False) + + mock_wmd.return_value.compute_batch.return_value = [ + ("t1", 0.5), + ("t2", np.inf), + ("t3", 1.5), + ] + mock_trainer.compute_weights.return_value = np.ones(3) + mock_surrogate = MagicMock() + mock_surrogate.coefficients.return_value = np.array([0.1, 0.2]) + mock_surrogate.predict.return_value = np.array([0.5, 0.6, 0.7]) + mock_factory.create.return_value = mock_surrogate + mock_metrics.calculate.return_value = MagicMock() + + result = explainer.explain("hello world") + + distances = result.raw_data["distances"] + assert distances[0] == pytest.approx(0.5) + assert distances[1] == pytest.approx(1001.5) + assert distances[2] == pytest.approx(1.5) + + +@patch("xwhy.explainers.text.SurrogateTrainer") +@patch("xwhy.explainers.text.SurrogateFactory") +@patch("xwhy.explainers.text.RegressionMetrics") +@patch("xwhy.explainers.text.WMDDistance") +def test_text_explain_impute_when_all_distances_non_finite( + mock_wmd: MagicMock, + mock_metrics: MagicMock, + mock_factory: MagicMock, + mock_trainer: MagicMock, +) -> None: + """Cover the branch where every WMD distance is non-finite.""" + predict_fn = MagicMock(return_value=np.array([[0.2, 0.8], [0.3, 0.7]])) + + with ( + patch("xwhy.explainers.text.EmbeddingFactory"), + patch("xwhy.explainers.text.TextPerturbation") as mock_pert, + ): + mock_pert.return_value.generate.return_value = ( + ["t1", "t2"], + [[1, 0], [0, 1]], + ) + explainer = TextExplainer(predict_fn=predict_fn, use_best_surrogate=False) + + mock_wmd.return_value.compute_batch.return_value = [ + ("t1", np.inf), + ("t2", np.nan), + ] + mock_trainer.compute_weights.return_value = np.ones(2) + mock_surrogate = MagicMock() + mock_surrogate.coefficients.return_value = np.array([0.1]) + mock_surrogate.predict.return_value = np.array([0.5, 0.5]) + mock_factory.create.return_value = mock_surrogate + mock_metrics.calculate.return_value = MagicMock() + + result = explainer.explain("hello world") + + distances = result.raw_data["distances"] + assert distances[0] == pytest.approx(1000.0) + assert distances[1] == pytest.approx(1000.0) diff --git a/tests/providers/test_anthropic.py b/tests/providers/test_anthropic.py index 90f1d6db..af958249 100644 --- a/tests/providers/test_anthropic.py +++ b/tests/providers/test_anthropic.py @@ -1,11 +1,17 @@ -"""Unit tests for the Anthropic provider functionality.""" +"""Tests for the Anthropic provider.""" -from unittest.mock import MagicMock +import re +from typing import Any +from unittest.mock import MagicMock, call, patch import pytest from xwhy.providers.anthropic import AnthropicProvider +# ------------------------------------------------------------------------- +# Text Generation Tests +# ------------------------------------------------------------------------- + def test_anthropic_provider_success() -> None: """Test successful text generation with Anthropic.""" @@ -36,36 +42,80 @@ def test_anthropic_provider_success() -> None: ) -def test_anthropic_provider_api_error() -> None: - """Test general exception handling during Anthropic API calls.""" +@patch("time.sleep", return_value=None) +def test_anthropic_provider_api_error_max_retries( + mock_sleep: MagicMock, +) -> None: + """Test generic exception handling during Anthropic API calls with retries.""" mock_client = MagicMock() - - # Simulate a network or authentication error mock_client.messages.create.side_effect = Exception("Invalid API Key or Limit") provider = AnthropicProvider(client=mock_client) - with pytest.raises(RuntimeError, match="Invalid API Key or Limit"): - provider.answer(prompt="Will fail") + with pytest.raises( + RuntimeError, match="Anthropic request failed: Invalid API Key or Limit" + ): + provider.answer(prompt="Will fail", max_retries=3) - mock_client.messages.create.assert_called_once_with( - model="claude-opus-4-8", - max_tokens=1024, - temperature=0.0, - messages=[{"role": "user", "content": "Will fail"}], - ) + assert mock_client.messages.create.call_count == 3 + assert mock_sleep.call_count == 2 -def test_anthropic_empty_response_content_raises_error() -> None: - """Test RuntimeError is raised when Anthropic returns empty content. - - This covers the scenario where response.content is an empty list, - leaving result_text empty and triggering the safety RuntimeError. - """ +@patch("time.sleep", return_value=None) +def test_anthropic_retry_then_success(mock_sleep: MagicMock) -> None: + """Test retry logic when API fails transiently before succeeding.""" mock_client = MagicMock() mock_response = MagicMock() + mock_content_block = MagicMock() + mock_content_block.text = "Success after retry" + mock_response.content = [mock_content_block] + + mock_client.messages.create.side_effect = [ + Exception("Temporary network glitch"), + mock_response, + ] - mock_response.content = [] + provider = AnthropicProvider(client=mock_client) + result = provider.answer(prompt="Test retry", max_retries=3, delay=1.0) + + assert result == "Success after retry" + assert mock_client.messages.create.call_count == 2 + mock_sleep.assert_called_once_with(1.0) + + +@patch("time.sleep", return_value=None) +def test_anthropic_direct_runtime_error_raises_immediately( + mock_sleep: MagicMock, +) -> None: + """RuntimeError raised during API execution should re-raise without retrying.""" + mock_client = MagicMock() + mock_client.messages.create.side_effect = RuntimeError("Direct RuntimeError") + + provider = AnthropicProvider(client=mock_client) + + with pytest.raises(RuntimeError, match="Direct RuntimeError"): + provider.answer(prompt="Will fail immediately", max_retries=3) + + assert mock_client.messages.create.call_count == 1 + mock_sleep.assert_not_called() + + +@pytest.mark.parametrize( + "empty_content", + [ + [], + [MagicMock(text=" \n ")], + ], +) +@patch("time.sleep", return_value=None) +def test_anthropic_empty_response_content_raises_error( + mock_sleep: MagicMock, + empty_content: Any, # noqa: ANN401 +) -> None: + """Test RuntimeError is raised immediately when Anthropic returns empty content.""" + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.content = empty_content mock_client.messages.create.return_value = mock_response provider = AnthropicProvider(client=mock_client) @@ -75,3 +125,46 @@ def test_anthropic_empty_response_content_raises_error() -> None: provider.answer(prompt="Test empty response") mock_client.messages.create.assert_called_once() + mock_sleep.assert_not_called() + + +@patch("time.sleep", return_value=None) +def test_anthropic_exponential_backoff(mock_sleep: MagicMock) -> None: + """Ensure retries respect exponential backoff capped at 30s.""" + mock_client = MagicMock() + mock_client.messages.create.side_effect = Exception("API Error") + + provider = AnthropicProvider(client=mock_client) + with pytest.raises(RuntimeError, match="Anthropic request failed"): + provider.answer(prompt="Backoff test", max_retries=6) + + expected_calls = [call(2), call(4), call(8), call(16), call(30)] + mock_sleep.assert_has_calls(expected_calls) + assert mock_sleep.call_count == 5 + + +@patch("time.sleep", return_value=None) +def test_anthropic_custom_delay(mock_sleep: MagicMock) -> None: + """Ensure custom delay overrides exponential backoff.""" + mock_client = MagicMock() + mock_client.messages.create.side_effect = Exception("API Error") + + provider = AnthropicProvider(client=mock_client) + with pytest.raises(RuntimeError, match="Anthropic request failed"): + provider.answer(prompt="Custom delay test", max_retries=3, delay=2.5) + + expected_calls = [call(2.5), call(2.5)] + mock_sleep.assert_has_calls(expected_calls) + assert mock_sleep.call_count == 2 + + +def test_anthropic_zero_retries_raises_fallback() -> None: + """Hit the end-of-function fallback RuntimeError by supplying max_retries=0.""" + mock_client = MagicMock() + provider = AnthropicProvider(mock_client) + + with pytest.raises( + RuntimeError, + match=re.escape("Anthropic text generation failed after max retries."), + ): + provider.answer("prompt", max_retries=0) diff --git a/tests/providers/test_gemini.py b/tests/providers/test_gemini.py index 91c8879b..176c0c0d 100644 --- a/tests/providers/test_gemini.py +++ b/tests/providers/test_gemini.py @@ -1,13 +1,17 @@ """Unit tests for the Gemini provider functionality.""" import json -from unittest.mock import MagicMock, PropertyMock, mock_open, patch +from unittest.mock import MagicMock, PropertyMock, call, mock_open, patch import pytest from PIL import Image from xwhy.providers.gemini import GeminiProvider +# ------------------------------------------------------------------------- +# Text Generation Tests +# ------------------------------------------------------------------------- + @patch("xwhy.providers.gemini.types") def test_gemini_provider_success(mock_types: MagicMock) -> None: @@ -44,8 +48,12 @@ def test_gemini_provider_success(mock_types: MagicMock) -> None: ) +@patch("time.sleep", return_value=None) @patch("xwhy.providers.gemini.types") -def test_gemini_provider_safety_block_fallback(mock_types: MagicMock) -> None: +def test_gemini_provider_safety_block_fallback( + mock_types: MagicMock, + mock_sleep: MagicMock, +) -> None: """Test that Gemini provider raises RuntimeError when blocked by safety filters.""" mock_client = MagicMock() mock_response = MagicMock() @@ -61,28 +69,39 @@ def test_gemini_provider_safety_block_fallback(mock_types: MagicMock) -> None: with pytest.raises( RuntimeError, match="blocked \\(likely due to safety filters\\)" ): - provider.answer(prompt="Blocked prompt test") + provider.answer(prompt="Blocked prompt test", max_retries=2) - mock_client.models.generate_content.assert_called_once() - mock_types.Part.from_text.assert_called_once_with(text="Blocked prompt test") + # Asserts retries happened and sleep was called once before failing on 2nd try + assert mock_client.models.generate_content.call_count == 2 + mock_sleep.assert_called_once() + mock_types.Part.from_text.assert_called_with(text="Blocked prompt test") +@patch("time.sleep", return_value=None) @patch("xwhy.providers.gemini.types") -def test_gemini_provider_api_error(mock_types: MagicMock) -> None: - """Test general exception handling during API calls.""" +def test_gemini_provider_api_error( + mock_types: MagicMock, + mock_sleep: MagicMock, +) -> None: + """Test general exception handling and retries during API calls.""" mock_client = MagicMock() mock_client.models.generate_content.side_effect = Exception("API error") provider = GeminiProvider(client=mock_client) with pytest.raises(RuntimeError, match="API error"): - provider.answer(prompt="Error prompt test") + provider.answer(prompt="Error prompt test", max_retries=3) - mock_client.models.generate_content.assert_called_once() + assert mock_client.models.generate_content.call_count == 3 + assert mock_sleep.call_count == 2 +@patch("time.sleep", return_value=None) @patch("xwhy.providers.gemini.types") -def test_gemini_empty_text_response_raises_error(mock_types: MagicMock) -> None: +def test_gemini_empty_text_response_raises_error( + mock_types: MagicMock, + mock_sleep: MagicMock, +) -> None: """Test RuntimeError is raised when Gemini returns empty text directly.""" mock_client = MagicMock() mock_response = MagicMock() @@ -94,16 +113,77 @@ def test_gemini_empty_text_response_raises_error(mock_types: MagicMock) -> None: expected_error = "empty response from the Gemini API" with pytest.raises(RuntimeError, match=expected_error): - provider.answer(prompt="Test empty response") + provider.answer(prompt="Test empty response", max_retries=2) + + assert mock_client.models.generate_content.call_count == 2 + mock_sleep.assert_called_once() + - mock_client.models.generate_content.assert_called_once() +def test_gemini_provider_zero_retries() -> None: + """Test text generation fails immediately if max_retries is less than 1.""" + provider = GeminiProvider(client=MagicMock()) + with pytest.raises(RuntimeError, match="max_retries must be at least 1"): + provider.answer(prompt="Test", max_retries=0) + +@patch("time.sleep", return_value=None) +@patch("xwhy.providers.gemini.types") +def test_gemini_generate_success_after_retries( + mock_types: MagicMock, + mock_sleep: MagicMock, +) -> None: + """Test generation succeeds on a subsequent retry with an explicit delay.""" + mock_client = MagicMock() + mock_response = MagicMock() + type(mock_response).text = PropertyMock(return_value="Delayed success") + + mock_client.models.generate_content.side_effect = [ + Exception("Temporary failure"), + mock_response, + ] + + provider = GeminiProvider(client=mock_client) + result = provider.answer(prompt="Test", max_retries=3, delay=5.5) + + assert result == "Delayed success" + assert mock_client.models.generate_content.call_count == 2 + mock_sleep.assert_called_once_with(5.5) + + +@patch("time.sleep", return_value=None) +@patch("xwhy.providers.gemini.types") +def test_gemini_generate_exponential_backoff_max( + mock_types: MagicMock, + mock_sleep: MagicMock, +) -> None: + """Test exponential backoff correctly caps at 30 seconds across retries.""" + mock_client = MagicMock() + mock_client.models.generate_content.side_effect = Exception("Fail") + + provider = GeminiProvider(client=mock_client) + + with pytest.raises(RuntimeError): + # 6 retries mean 5 sleeps: 2, 4, 8, 16, 30 (cap) + provider.answer(prompt="Test", max_retries=6) + + expected_sleep_calls = [call(2), call(4), call(8), call(16), call(30)] + mock_sleep.assert_has_calls(expected_sleep_calls) + assert mock_sleep.call_count == 5 + + +# ------------------------------------------------------------------------- +# Image Generation & Execution Tests +# ------------------------------------------------------------------------- + + +@patch("time.sleep", return_value=None) @patch("xwhy.providers.gemini.Image.open") @patch("os.makedirs") def test_generate_image_stream_no_inline_data( mock_makedirs: MagicMock, mock_image_open: MagicMock, + mock_sleep: MagicMock, ) -> None: """Test stream generation where chunks/parts lack inline_data.""" client = MagicMock() @@ -117,18 +197,21 @@ def test_generate_image_stream_no_inline_data( provider._create_placeholder_image = MagicMock(return_value=None) # type: ignore[method-assign] success, _ = provider.generate_image( - prompt="Test", output_dir="fake_dir", stream=True + prompt="Test", output_dir="fake_dir", stream=True, max_retries=2 ) assert success is False + assert mock_sleep.call_count == 1 provider._create_placeholder_image.assert_called_once() +@patch("time.sleep", return_value=None) @patch("xwhy.providers.gemini.Image.open") @patch("os.makedirs") def test_generate_image_no_stream_no_inline_data( mock_makedirs: MagicMock, mock_image_open: MagicMock, + mock_sleep: MagicMock, ) -> None: """Test non-stream generation where response parts lack inline_data.""" client = MagicMock() @@ -142,16 +225,19 @@ def test_generate_image_no_stream_no_inline_data( provider._create_placeholder_image = MagicMock(return_value=None) # type: ignore[method-assign] success, _ = provider.generate_image( - prompt="Test", output_dir="fake_dir", stream=False + prompt="Test", output_dir="fake_dir", stream=False, max_retries=2 ) assert success is False + assert mock_sleep.call_count == 1 provider._create_placeholder_image.assert_called_once() +@patch("time.sleep", return_value=None) @patch("os.makedirs") def test_execute_image_request_fallback_not_pil_image( mock_makedirs: MagicMock, + mock_sleep: MagicMock, ) -> None: """Test when fallback image is not a PIL Image instance.""" client = MagicMock() @@ -160,17 +246,13 @@ def test_execute_image_request_fallback_not_pil_image( provider = GeminiProvider(client) provider._create_placeholder_image = MagicMock(return_value=12345) # type: ignore[method-assign] - success, _ = provider.generate_image("Test", "out", stream=False) + success, _ = provider.generate_image("Test", "out", stream=False, max_retries=1) assert success is False provider._create_placeholder_image.assert_called_once() -# ------------------------------------------------------------------------- -# Image Generation Tests -# ------------------------------------------------------------------------- - - +@patch("time.sleep", return_value=None) @patch("xwhy.providers.gemini.Image.open") @patch("os.makedirs") @patch("time.time", return_value=1234567.89) @@ -178,6 +260,7 @@ def test_generate_image_stream_success( mock_time: MagicMock, mock_makedirs: MagicMock, mock_image_open: MagicMock, + mock_sleep: MagicMock, ) -> None: """Test successful image generation with stream enabled.""" client = MagicMock() @@ -201,11 +284,13 @@ def test_generate_image_stream_success( mock_img_instance.save.assert_called_once() +@patch("time.sleep", return_value=None) @patch("xwhy.providers.gemini.Image.open") @patch("os.makedirs") def test_generate_image_no_stream_success( mock_makedirs: MagicMock, mock_image_open: MagicMock, + mock_sleep: MagicMock, ) -> None: """Test successful image generation with stream disabled.""" client = MagicMock() @@ -229,23 +314,57 @@ def test_generate_image_no_stream_success( mock_img_instance.save.assert_called_once() +@patch("time.sleep", return_value=None) @patch("os.makedirs") -def test_generate_image_exception_fallback(mock_makedirs: MagicMock) -> None: +def test_generate_image_exception_fallback( + mock_makedirs: MagicMock, + mock_sleep: MagicMock, +) -> None: """Test exception during API call triggers the placeholder fallback logic.""" client = MagicMock() client.models.generate_content_stream.side_effect = Exception("Stream fail") provider = GeminiProvider(client) - - # Mock fallback to return a non-image (string) to test type checking branch provider._create_placeholder_image = MagicMock(return_value="not_an_image") # type: ignore[method-assign] - success, _ = provider.generate_image("Test", "out", stream=True) + success, _ = provider.generate_image("Test", "out", stream=True, max_retries=2) assert success is False + assert mock_sleep.call_count == 1 provider._create_placeholder_image.assert_called_once() +@patch("time.sleep", return_value=None) +@patch("xwhy.providers.gemini.Image.open") +@patch("os.makedirs") +def test_generate_image_success_after_retry_with_delay( + mock_makedirs: MagicMock, + mock_image_open: MagicMock, + mock_sleep: MagicMock, +) -> None: + """Test image generation success on a subsequent retry with an explicit delay.""" + client = MagicMock() + + mock_response = MagicMock() + mock_part = MagicMock() + mock_part.inline_data = MagicMock(data=b"img", mime_type="image/png") + mock_response.parts = [mock_part] + + client.models.generate_content.side_effect = [Exception("Fail"), mock_response] + + mock_img_instance = MagicMock(spec=Image.Image) + mock_image_open.return_value = mock_img_instance + + provider = GeminiProvider(client) + success, _ = provider.generate_image( + prompt="Test", output_dir="out", stream=False, max_retries=2, delay=4.2 + ) + + assert success is True + mock_sleep.assert_called_once_with(4.2) + assert client.models.generate_content.call_count == 2 + + # ------------------------------------------------------------------------- # Image Editing Tests # ------------------------------------------------------------------------- @@ -328,7 +447,7 @@ def test_edit_image_png_success( @patch("os.remove") @patch("builtins.open") def test_submit_image_batch_with_image_and_seed( - mock_open: MagicMock, mock_remove: MagicMock + mock_open_func: MagicMock, mock_remove: MagicMock ) -> None: """Test batch submission with a base image and a deterministic seed.""" client = MagicMock() @@ -343,21 +462,20 @@ def test_submit_image_batch_with_image_and_seed( ) assert job_name == str(mock_job.name) - assert client.files.upload.call_count == 2 # Once for img, once for JSONL + assert client.files.upload.call_count == 2 mock_remove.assert_called_once() @patch("os.remove") @patch("builtins.open") def test_submit_image_batch_no_image_os_error( - mock_open: MagicMock, mock_remove: MagicMock + mock_open_func: MagicMock, mock_remove: MagicMock ) -> None: - """Test batch submission without a base image, and catching OSError on cleanup.""" + """Test batch submission without base image, catching OSError on cleanup.""" client = MagicMock() mock_job = MagicMock(name="job_123") client.batches.create.return_value = mock_job - # Trigger OSError in the try/except block for os.remove mock_remove.side_effect = OSError("Permission denied") provider = GeminiProvider(client) @@ -366,13 +484,13 @@ def test_submit_image_batch_no_image_os_error( ) assert job_name == str(mock_job.name) - client.files.upload.assert_called_once() # Only JSONL uploaded + client.files.upload.assert_called_once() @patch("os.makedirs") @patch("builtins.open") def test_retrieve_image_batch_success_found_image( - mock_open: MagicMock, mock_makedirs: MagicMock + mock_open_func: MagicMock, mock_makedirs: MagicMock ) -> None: """Test polling success and processing of JSONL with valid PNG and JPG data.""" client = MagicMock() @@ -380,7 +498,6 @@ def test_retrieve_image_batch_success_found_image( mock_job.state.name = "JOB_STATE_SUCCEEDED" client.batches.get.return_value = mock_job - # Construct JSONL with one PNG candidate and one JPG candidate lines = [ json.dumps( { @@ -405,7 +522,7 @@ def test_retrieve_image_batch_success_found_image( ), json.dumps( { - "key": "gemini_request_1_image", # Testing fallback "key" lookup + "key": "gemini_request_1_image", "response": { "candidates": [ { @@ -447,7 +564,6 @@ def test_retrieve_image_batch_no_inline_data_and_missing_result( mock_job.state.name = "JOB_STATE_SUCCEEDED" client.batches.get.return_value = mock_job - # Line 1: Valid JSON but missing inlineData. Line 2 (for t2) is completely missing. lines = [ json.dumps( { @@ -461,14 +577,10 @@ def test_retrieve_image_batch_no_inline_data_and_missing_result( client.files.download.return_value = b"\n".join(x.encode() for x in lines) provider = GeminiProvider(client) - - # Mock placeholder to test returning both a string path and None provider._create_placeholder_image = MagicMock(side_effect=["out.png", None]) # type: ignore[method-assign] results = provider.retrieve_image_batch(job_name="job1", text_list=["t1", "t2"]) - # First one appends because placeholder returns a string - # Second one skips appending because placeholder returns None assert len(results) == 1 assert results[0][0] is False assert results[0][1] == "out.png" @@ -484,7 +596,6 @@ def test_retrieve_image_batch_json_parse_error( mock_job.state.name = "JOB_STATE_SUCCEEDED" client.batches.get.return_value = mock_job - # One line with a parseable custom_id, one completely malformed content = ( b'{"custom_id": "gemini_request_0_image", invalid\n' b"completely bad string format\n" @@ -545,11 +656,13 @@ def test_retrieve_image_batch_empty_lines_and_non_string_placeholder( provider._create_placeholder_image.assert_called_once() +@patch("time.sleep", return_value=None) @patch("xwhy.providers.gemini.Image.open") @patch("os.makedirs") def test_execute_image_request_fallback_is_pil_image( mock_makedirs: MagicMock, mock_image_open: MagicMock, + mock_sleep: MagicMock, ) -> None: """Test when API fails and fallback image returns a valid PIL Image instance.""" client = MagicMock() @@ -559,7 +672,7 @@ def test_execute_image_request_fallback_is_pil_image( dummy_img = MagicMock(spec=Image.Image) provider._create_placeholder_image = MagicMock(return_value=dummy_img) # type: ignore[method-assign] - success, _ = provider.generate_image("Test", "out", stream=False) + success, _ = provider.generate_image("Test", "out", stream=False, max_retries=1) assert success is False dummy_img.save.assert_called_once() @@ -603,8 +716,6 @@ def test_retrieve_image_batch_comprehensive_coverage( client.batches.get.return_value = mock_job lines = [ - # 1. Valid JPEG inlineData (hits image/jpeg, .jpg extension, path is - # not None branch) json.dumps( { "custom_id": "gemini_request_0_image", @@ -626,8 +737,6 @@ def test_retrieve_image_batch_comprehensive_coverage( }, } ), - # 2. Candidate with no inlineData (hits "if not found_image:" and path is - # None -> placeholder) json.dumps( { "custom_id": "gemini_request_1_image", @@ -644,8 +753,6 @@ def test_retrieve_image_batch_comprehensive_coverage( provider = GeminiProvider(client) provider._create_placeholder_image = MagicMock(return_value="placeholder.png") # type: ignore[method-assign] - # Three items: index 0 (found), index 1 (not found -> path is None), - # index 2 (missing entirely from results) results = provider.retrieve_image_batch( job_name="job1", text_list=["prompt0", "prompt1", "prompt2"] ) @@ -761,7 +868,6 @@ def test_retrieve_image_batch_json_parse_error_without_custom_id( mock_job.dest.file_name = "results.jsonl" client.batches.get.return_value = mock_job - # Malformed line entirely missing "custom_id" to trigger the 'else' branch content = b"completely malformed line without custom id key\n" client.files.download.return_value = content @@ -778,12 +884,7 @@ def test_retrieve_image_batch_json_parse_error_without_custom_id( def test_retrieve_image_batch_path_none_non_string_placeholder( mock_makedirs: MagicMock, ) -> None: - """Cover path is None with non-string placeholder (isinstance false branch). - - When a custom_id is present in processed_results with path=None (API - refusal / no inlineData) and _create_placeholder_image returns a - non-string value, the result must not be appended to final_output_list. - """ + """Cover path is None with non-string placeholder.""" client = MagicMock() mock_job = MagicMock() mock_job.state.name = "JOB_STATE_SUCCEEDED" diff --git a/tests/providers/test_huggingface.py b/tests/providers/test_huggingface.py index 9bdc61fe..514a586c 100644 --- a/tests/providers/test_huggingface.py +++ b/tests/providers/test_huggingface.py @@ -1,8 +1,9 @@ """Unit tests for the HuggingFace provider functionality.""" +import re from collections.abc import Callable from typing import Any -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, call, patch import pytest import torch @@ -59,8 +60,9 @@ def test_huggingface_provider_success() -> None: ) -def test_huggingface_provider_api_error() -> None: - """Test general exception handling during HuggingFace API calls.""" +@patch("time.sleep", return_value=None) +def test_huggingface_provider_api_error(mock_sleep: MagicMock) -> None: + """Test general exception handling and retries during HuggingFace API calls.""" mock_client = MagicMock() mock_client.chat.completions.create.side_effect = Exception( "Model loading or API error" @@ -69,21 +71,17 @@ def test_huggingface_provider_api_error() -> None: provider = HuggingFaceProvider(client=mock_client) with pytest.raises(RuntimeError, match="Model loading or API error"): - provider.answer(prompt="Error prompt test") - - mock_client.chat.completions.create.assert_called_once_with( - model="meta-llama/Meta-Llama-3-8B-Instruct", - messages=[{"role": "user", "content": "Error prompt test"}], - max_tokens=512, - temperature=0.1, - ) + provider.answer(prompt="Error prompt test", max_retries=3) + assert mock_client.chat.completions.create.call_count == 3 + assert mock_sleep.call_count == 2 -def test_huggingface_empty_text_response_raises_error() -> None: - """Test RuntimeError is raised when HuggingFace returns empty text. - This covers the 'if not result_text:' block for the HuggingFace API. - """ +@patch("time.sleep", return_value=None) +def test_huggingface_empty_text_response_raises_error( + mock_sleep: MagicMock, +) -> None: + """Test RuntimeError is raised when HuggingFace returns empty text.""" mock_client = MagicMock() mock_response = MagicMock() mock_choice = MagicMock() @@ -96,9 +94,69 @@ def test_huggingface_empty_text_response_raises_error() -> None: expected_error = "empty response from the HuggingFace API" with pytest.raises(RuntimeError, match=expected_error): - provider.answer(prompt="Test empty response") + provider.answer(prompt="Test empty response", max_retries=2) + + assert mock_client.chat.completions.create.call_count == 2 + mock_sleep.assert_called_once() + + +@patch("time.sleep", return_value=None) +def test_huggingface_generate_retry_success(mock_sleep: MagicMock) -> None: + """Test text generation succeeds after retry with explicit delay.""" + mock_client = MagicMock() + mock_response = MagicMock() + mock_choice = MagicMock() + mock_choice.message.content = "Recovered text" + mock_response.choices = [mock_choice] + + mock_client.chat.completions.create.side_effect = [ + Exception("Temporary failure"), + mock_response, + ] + + provider = HuggingFaceProvider(client=mock_client) + result = provider.answer(prompt="Retry test", max_retries=3, delay=2.5) + + assert result == "Recovered text" + assert mock_client.chat.completions.create.call_count == 2 + mock_sleep.assert_called_once_with(2.5) + + +@patch("time.sleep", return_value=None) +def test_huggingface_generate_exponential_backoff( + mock_sleep: MagicMock, +) -> None: + """Test exponential backoff for text generation retries up to the 30s cap.""" + mock_client = MagicMock() + mock_client.chat.completions.create.side_effect = Exception("API error") + + provider = HuggingFaceProvider(client=mock_client) - mock_client.chat.completions.create.assert_called_once() + with pytest.raises(RuntimeError, match="HuggingFace request failed"): + provider.answer(prompt="Test backoff", max_retries=6) + + expected_sleep_calls = [call(2), call(4), call(8), call(16), call(30)] + mock_sleep.assert_has_calls(expected_sleep_calls) + assert mock_sleep.call_count == 5 + + +def test_huggingface_generate_zero_retries_raises_fallback() -> None: + """Test that zero retries triggers the fallback RuntimeError. + + The loop never executes when max_retries is set to zero, so control + falls through to the final exception raise. + + """ + mock_client = MagicMock() + provider = HuggingFaceProvider(client=mock_client) + + with pytest.raises( + RuntimeError, + match=re.escape("HuggingFace text generation failed after max retries."), + ): + provider.answer(prompt="prompt", max_retries=0) + + mock_client.chat.completions.create.assert_not_called() # --------------------------------------------------------------------------- @@ -214,10 +272,38 @@ def test_initialize_pipeline_instruct_pix2pix() -> None: ) mock_from_pretrained.assert_called_once() + assert mock_from_pretrained.call_args.kwargs["dtype"] == torch.float32 mock_pipe.to.assert_called_once() assert provider.pipe is mock_pipe +def test_initialize_pipeline_instruct_pix2pix_cuda() -> None: + """Build an InstructPix2Pix pipeline when CUDA is available (float16).""" + mock_client = MagicMock() + mock_pipe = MagicMock() + mock_pipe.scheduler.config = {} + + with ( + patch("torch.cuda.is_available", return_value=True), + patch( + "diffusers.StableDiffusionInstructPix2PixPipeline.from_pretrained", + return_value=mock_pipe, + ) as mock_from_pretrained, + patch( + "diffusers.EulerAncestralDiscreteScheduler.from_config", + return_value=MagicMock(), + ), + ): + provider = HuggingFaceProvider( + client=mock_client, + model_name="timbrooks/instruct-pix2pix", + ) + + mock_from_pretrained.assert_called_once() + assert mock_from_pretrained.call_args.kwargs["dtype"] == torch.float16 + assert provider.pipe is mock_pipe + + def test_initialize_pipeline_inpaint_requires_segmentation() -> None: """Raise RuntimeError when an inpaint model is requested without segmentation.""" mock_client = MagicMock() @@ -252,10 +338,34 @@ def test_initialize_pipeline_inpaint_success() -> None: ) mock_from_pretrained.assert_called_once() + assert mock_from_pretrained.call_args.kwargs["dtype"] == torch.float32 mock_pipe.to.assert_called_once() assert provider.pipe is mock_pipe +def test_initialize_pipeline_inpaint_cuda() -> None: + """Build an inpainting pipeline when CUDA is available (float16).""" + mock_client = MagicMock() + mock_pipe = MagicMock() + + with ( + patch("torch.cuda.is_available", return_value=True), + patch( + "diffusers.StableDiffusionInpaintPipeline.from_pretrained", + return_value=mock_pipe, + ) as mock_from_pretrained, + ): + provider = HuggingFaceProvider( + client=mock_client, + model_name="runwayml/stable-diffusion-inpainting", + use_segmentation_model=True, + ) + + mock_from_pretrained.assert_called_once() + assert mock_from_pretrained.call_args.kwargs["dtype"] == torch.float16 + assert provider.pipe is mock_pipe + + def test_initialize_pipeline_auto_text2image_success() -> None: """Use AutoPipelineForText2Image for a generic diffusers model name.""" mock_client = MagicMock() @@ -274,10 +384,33 @@ def test_initialize_pipeline_auto_text2image_success() -> None: ) mock_from_pretrained.assert_called_once() + assert mock_from_pretrained.call_args.kwargs["torch_dtype"] == torch.float32 mock_pipe.to.assert_called_once() assert provider.pipe is mock_pipe +def test_initialize_pipeline_auto_text2image_cuda() -> None: + """Use AutoPipelineForText2Image with CUDA enabled (float16).""" + mock_client = MagicMock() + mock_pipe = MagicMock() + + with ( + patch("torch.cuda.is_available", return_value=True), + patch( + "diffusers.AutoPipelineForText2Image.from_pretrained", + return_value=mock_pipe, + ) as mock_from_pretrained, + ): + provider = HuggingFaceProvider( + client=mock_client, + model_name="stabilityai/stable-diffusion-2", + ) + + mock_from_pretrained.assert_called_once() + assert mock_from_pretrained.call_args.kwargs["torch_dtype"] == torch.float16 + assert provider.pipe is mock_pipe + + def test_initialize_pipeline_falls_back_to_diffusion_pipeline() -> None: """Fall back to DiffusionPipeline when AutoPipelineForText2Image fails.""" mock_client = MagicMock() @@ -300,6 +433,33 @@ def test_initialize_pipeline_falls_back_to_diffusion_pipeline() -> None: ) mock_diffusion.assert_called_once() + assert mock_diffusion.call_args.kwargs["torch_dtype"] == torch.float32 + assert provider.pipe is mock_pipe + + +def test_initialize_pipeline_falls_back_to_diffusion_pipeline_cuda() -> None: + """Fall back to DiffusionPipeline when CUDA is available (float16).""" + mock_client = MagicMock() + mock_pipe = MagicMock() + + with ( + patch("torch.cuda.is_available", return_value=True), + patch( + "diffusers.AutoPipelineForText2Image.from_pretrained", + side_effect=Exception("not a text2image model"), + ), + patch( + "diffusers.DiffusionPipeline.from_pretrained", + return_value=mock_pipe, + ) as mock_diffusion, + ): + provider = HuggingFaceProvider( + client=mock_client, + model_name="some/other-model", + ) + + mock_diffusion.assert_called_once() + assert mock_diffusion.call_args.kwargs["torch_dtype"] == torch.float16 assert provider.pipe is mock_pipe @@ -371,11 +531,7 @@ def test_execute_image_request_raises_when_pipe_none() -> None: def test_generate_image_success(tmp_path: Any) -> None: # noqa: ANN401 - """Successfully generate an image and write it to disk. - - Explicitly supplies num_inference_steps so that the - ``if "num_inference_steps" not in kwargs`` branch evaluates to False. - """ + """Successfully generate an image and write it to disk.""" mock_pipe = _make_pipe("StableDiffusionPipeline") mock_image = MagicMock(spec=Image.Image) mock_pipe.return_value = MagicMock(images=[mock_image]) @@ -384,7 +540,7 @@ def test_generate_image_success(tmp_path: Any) -> None: # noqa: ANN401 success, path = provider.generate_image( prompt="a red cube", output_dir=str(tmp_path), - num_inference_steps=20, # already present → if-branch skipped + num_inference_steps=20, ) assert success is True @@ -411,10 +567,12 @@ def test_generate_image_output_as_list(tmp_path: Any) -> None: # noqa: ANN401 assert path.endswith(".png") +@patch("time.sleep", return_value=None) def test_generate_image_no_valid_output_uses_fallback( + mock_sleep: MagicMock, tmp_path: Any, # noqa: ANN401 ) -> None: - """Fall back to a placeholder image when the pipeline returns nothing useful.""" + """Fall back to a placeholder image when pipeline returns no valid image.""" mock_pipe = _make_pipe("StableDiffusionPipeline") mock_pipe.return_value = MagicMock(images=[]) @@ -429,11 +587,40 @@ def test_generate_image_no_valid_output_uses_fallback( success, path = provider.generate_image( prompt="broken", output_dir=str(tmp_path), + max_retries=2, ) assert success is False placeholder.save.assert_called_once() assert path.endswith(".png") + mock_sleep.assert_called_once() + + +@patch("time.sleep", return_value=None) +def test_generate_image_retry_success( + mock_sleep: MagicMock, + tmp_path: Any, # noqa: ANN401 +) -> None: + """Test image generation retry success on second attempt with explicit delay.""" + mock_pipe = _make_pipe("StableDiffusionPipeline") + mock_image = MagicMock(spec=Image.Image) + mock_pipe.side_effect = [ + Exception("Temporary failure"), + MagicMock(images=[mock_image]), + ] + + provider = HuggingFaceProvider(client=MagicMock(), pipe=mock_pipe) + success, path = provider.generate_image( + prompt="a cat", + output_dir=str(tmp_path), + max_retries=3, + delay=3.5, + ) + + assert success is True + assert path.endswith(".png") + mock_sleep.assert_called_once_with(3.5) + assert mock_pipe.call_count == 2 def test_edit_image_file_not_found() -> None: @@ -458,7 +645,7 @@ def test_edit_image_success( ) -> None: """Edit an existing image with a generic (non-inpaint) pipeline.""" src_path = tmp_path / "src.png" - src_path.write_bytes(b"") # satisfy os.path.exists + src_path.write_bytes(b"") mock_pil = MagicMock(spec=Image.Image) mock_pil.convert.return_value = mock_pil @@ -541,9 +728,9 @@ def test_edit_image_pix2pix_sets_guidance( mock_open: MagicMock, tmp_path: Any, # noqa: ANN401 ) -> None: - """Inject the default image_guidance_scale for InstructPix2Pix pipelines.""" + """Inject default image_guidance_scale for InstructPix2Pix pipelines.""" src_path = tmp_path / "src.png" - src_path.write_bytes(b"") # satisfy os.path.exists + src_path.write_bytes(b"") mock_pil = MagicMock(spec=Image.Image) mock_pil.convert.return_value = mock_pil @@ -566,7 +753,9 @@ def test_edit_image_pix2pix_sets_guidance( assert call_kwargs.get("image_guidance_scale") == 1.0 +@patch("time.sleep", return_value=None) def test_execute_image_request_pipeline_exception_uses_fallback( + mock_sleep: MagicMock, tmp_path: Any, # noqa: ANN401 ) -> None: """Catch pipeline exceptions, set success=False and save a placeholder.""" @@ -584,11 +773,13 @@ def test_execute_image_request_pipeline_exception_uses_fallback( success, path = provider.generate_image( prompt="boom", output_dir=str(tmp_path), + max_retries=2, ) assert success is False placeholder.save.assert_called_once() assert path.endswith(".png") + mock_sleep.assert_called_once() @patch("xwhy.providers.huggingface.Image.open") @@ -598,14 +789,7 @@ def test_edit_image_inpaint_without_segmentation_raises( mock_open: MagicMock, tmp_path: Any, # noqa: ANN401 ) -> None: - """Raise ValueError when an inpaint pipe is used without a segmentation model. - - Covers the exact branch: - elif is_inpaint and input_image_path is not None: - raise ValueError( - "segmentation_model is required for Inpainting pipelines." - ) - """ + """Raise ValueError when inpaint pipe is used without segmentation model.""" src_path = tmp_path / "src.png" src_path.write_bytes(b"") @@ -626,19 +810,13 @@ def test_edit_image_inpaint_without_segmentation_raises( prompt="fill", image_path=str(src_path), output_dir=str(tmp_path), - # segmentation_model deliberately omitted ) def test_generate_image_default_inference_steps( tmp_path: Any, # noqa: ANN401 ) -> None: - """Inject num_inference_steps=30 for non-inpaint pipelines. - - Covers the True arm of: - if "num_inference_steps" not in kwargs: - kwargs["num_inference_steps"] = 50 if is_inpaint else 30 - """ + """Inject num_inference_steps=30 for non-inpaint pipelines.""" mock_pipe = _make_pipe("StableDiffusionPipeline") mock_image = MagicMock(spec=Image.Image) mock_pipe.return_value = MagicMock(images=[mock_image]) @@ -658,11 +836,7 @@ def test_edit_image_inpaint_default_inference_steps( mock_get_mask: MagicMock, tmp_path: Any, # noqa: ANN401 ) -> None: - """Inject num_inference_steps=50 for inpaint pipelines. - - Covers the True arm of the same if-statement and the True arm of the - ternary ``50 if is_inpaint else 30``. - """ + """Inject num_inference_steps=50 for inpaint pipelines.""" src_path = tmp_path / "src.png" src_path.write_bytes(b"") @@ -691,10 +865,12 @@ def test_edit_image_inpaint_default_inference_steps( assert mock_pipe.call_args.kwargs["num_inference_steps"] == 50 +@patch("time.sleep", return_value=None) def test_pipeline_exception_fallback_is_image( + mock_sleep: MagicMock, tmp_path: Any, # noqa: ANN401 ) -> None: - """Assign the placeholder when _create_placeholder_image returns an Image.""" + """Assign placeholder when _create_placeholder_image returns an Image.""" mock_pipe = _make_pipe("StableDiffusionPipeline") mock_pipe.side_effect = RuntimeError("boom") @@ -702,17 +878,21 @@ def test_pipeline_exception_fallback_is_image( provider = HuggingFaceProvider(client=MagicMock(), pipe=mock_pipe) with patch.object(provider, "_create_placeholder_image", return_value=placeholder): - success, path = provider.generate_image(prompt="x", output_dir=str(tmp_path)) + success, path = provider.generate_image( + prompt="x", output_dir=str(tmp_path), max_retries=1 + ) assert success is False placeholder.save.assert_called_once() assert path.endswith(".png") +@patch("time.sleep", return_value=None) def test_pipeline_exception_fallback_not_image( + mock_sleep: MagicMock, tmp_path: Any, # noqa: ANN401 ) -> None: - """Skip save when the fallback is not a PIL Image.""" + """Skip save when fallback is not a PIL Image.""" mock_pipe = _make_pipe("StableDiffusionPipeline") mock_pipe.side_effect = RuntimeError("boom") @@ -721,10 +901,11 @@ def test_pipeline_exception_fallback_not_image( with patch.object( provider, "_create_placeholder_image", return_value="not-an-image" ): - success, path = provider.generate_image(prompt="x", output_dir=str(tmp_path)) + success, path = provider.generate_image( + prompt="x", output_dir=str(tmp_path), max_retries=1 + ) assert success is False - # path is still built, but no .save() occurred assert path.endswith(".png") @@ -751,3 +932,43 @@ def test_generate_image_inpaint_pipe_no_mask_required( assert success is True assert path.endswith(".png") mock_image.save.assert_called_once() + + +@patch("time.sleep", return_value=None) +def test_execute_image_request_generated_img_none_no_break( + mock_sleep: MagicMock, + tmp_path: Any, # noqa: ANN401 +) -> None: + """Test retry continuation when generated image is None. + + Verifies that if the pipeline returns a list where the first element + is None, the retry loop continues instead of breaking. + + Args: + mock_sleep: Mock for time.sleep to avoid real delays. + tmp_path: Pytest temporary directory fixture. + + """ + mock_pipe = _make_pipe("StableDiffusionPipeline") + mock_pipe.return_value = MagicMock(images=[None]) + + placeholder = MagicMock(spec=Image.Image) + provider = HuggingFaceProvider(client=MagicMock(), pipe=mock_pipe) + + with patch.object( + provider, + "_create_placeholder_image", + return_value=placeholder, + ): + success, path = provider.generate_image( + prompt="none image", + output_dir=str(tmp_path), + max_retries=2, + ) + + assert success is False + assert path.endswith(".png") + # Both attempts ran (no break occurred). + assert mock_pipe.call_count == 2 + mock_sleep.assert_called_once() + placeholder.save.assert_called_once() diff --git a/tests/providers/test_openai.py b/tests/providers/test_openai.py index 7c635d59..bc100d7c 100644 --- a/tests/providers/test_openai.py +++ b/tests/providers/test_openai.py @@ -1,28 +1,30 @@ """Tests for the OpenAI provider.""" -from unittest.mock import MagicMock, PropertyMock, patch +import re +from typing import Any +from unittest.mock import MagicMock, PropertyMock, call, patch import pytest from PIL import Image from xwhy.providers.openai import OpenAIProvider +# ------------------------------------------------------------------------- +# Text Generation & Reasoning Tests +# ------------------------------------------------------------------------- + def test_answer_uses_completion_api() -> None: """Completion models should use the Completions API.""" client = MagicMock() - response = MagicMock() response.choices = [MagicMock(text="hello")] - client.completions.create.return_value = response provider = OpenAIProvider(client) - result = provider.answer("prompt") assert result == "hello" - client.completions.create.assert_called_once() client.responses.create.assert_not_called() @@ -30,46 +32,42 @@ def test_answer_uses_completion_api() -> None: def test_answer_uses_responses_api() -> None: """Reasoning models should use the Responses API.""" client = MagicMock() - response = MagicMock() response.output_text = "reasoning" - client.responses.create.return_value = response provider = OpenAIProvider(client) - - result = provider.answer( - "prompt", - model="gpt-5-mini", - ) + result = provider.answer("prompt", model="gpt-5-mini") assert result == "reasoning" - client.responses.create.assert_called_once() client.completions.create.assert_not_called() def test_answer_raises_runtime_error_when_client_fails() -> None: - """RuntimeError from the client should propagate and raise RuntimeError.""" + """RuntimeError from the client should propagate directly.""" client = MagicMock() - client.completions.create.side_effect = RuntimeError("boom") provider = OpenAIProvider(client) - with pytest.raises(RuntimeError, match="boom"): provider.answer("prompt") -def test_answer_raises_runtime_error_on_generic_exception() -> None: +@patch("time.sleep", return_value=None) +def test_answer_raises_runtime_error_on_generic_exception( + mock_sleep: MagicMock, +) -> None: """Generic exceptions should be caught, logged, and raise a RuntimeError.""" client = MagicMock() client.completions.create.side_effect = ValueError("generic error") provider = OpenAIProvider(client) - with pytest.raises(RuntimeError, match="OpenAI request failed: generic error"): - provider.answer("prompt") + provider.answer("prompt", max_retries=3) + + assert client.completions.create.call_count == 3 + assert mock_sleep.call_count == 2 @pytest.mark.parametrize( @@ -84,10 +82,7 @@ def test_answer_raises_runtime_error_on_generic_exception() -> None: ("o4-mini", True), ], ) -def test_is_reasoning_model( - model: str, - expected: bool, -) -> None: +def test_is_reasoning_model(model: str, expected: bool) -> None: """Reasoning models should be correctly detected.""" assert OpenAIProvider._is_reasoning_model(model) is expected @@ -111,7 +106,6 @@ def test_generate_regex_dynamic_fix() -> None: ) assert result == "fixed_response" - assert client.completions.create.call_count == 2 retry_call = client.completions.create.call_args_list[1] @@ -121,14 +115,12 @@ def test_generate_regex_dynamic_fix() -> None: def test_generate_regex_no_match_fallback() -> None: """Raise RuntimeError when regex matching fails on tokens error.""" client = MagicMock() - error_message = ( "Error: max_output_tokens is an integer below minimum value. Expected a value." ) client.completions.create.side_effect = Exception(error_message) provider = OpenAIProvider(client) - with patch("xwhy.providers.openai.logger") as mock_logger: with pytest.raises(RuntimeError, match="Expected a value"): provider._generate( @@ -136,11 +128,11 @@ def test_generate_regex_no_match_fallback() -> None: model="gpt-3.5-turbo-instruct", max_tokens=10, temperature=0.0, + max_retries=1, ) assert client.completions.create.call_count == 1 mock_logger.error.assert_called() - mock_logger.warning.assert_not_called() def test_openai_provider_reasoning_model_with_temperature() -> None: @@ -169,7 +161,6 @@ def test_openai_provider_reasoning_model_temperature_fallback() -> None: mock_response = MagicMock() mock_response.output_text = "Fallback success" - # First call raises temperature error, second call succeeds with temperature=1.0 mock_client.responses.create.side_effect = [ Exception("The temperature parameter is not supported with this model."), mock_response, @@ -180,10 +171,12 @@ def test_openai_provider_reasoning_model_temperature_fallback() -> None: assert result == "Fallback success" assert mock_client.responses.create.call_count == 2 + retry_call = mock_client.responses.create.call_args_list[1] + assert retry_call.kwargs["temperature"] == 1.0 def test_openai_provider_max_tokens_lowercase_regex() -> None: - """Test that token limitation errors are handled with the new lowercase regex.""" + """Test that token limitation errors are handled with the lowercase regex.""" mock_client = MagicMock() mock_response = MagicMock() mock_response.output_text = "Token fix success" @@ -210,9 +203,10 @@ def test_openai_provider_temperature_already_one_no_retry() -> None: ) provider = OpenAIProvider(client=mock_client) - with pytest.raises(RuntimeError, match="temperature parameter is not supported"): - provider.answer(prompt="Test", model="o1-preview", temperature=1.0) + provider.answer( + prompt="Test", model="o1-preview", temperature=1.0, max_retries=1 + ) assert mock_client.responses.create.call_count == 1 @@ -225,20 +219,19 @@ def test_openai_provider_max_tokens_no_regex_match() -> None: ) provider = OpenAIProvider(client=mock_client) - with pytest.raises(RuntimeError, match="unexpected error format"): - provider.answer(prompt="Test", model="o3-mini", max_tokens=5) + provider.answer(prompt="Test", model="o3-mini", max_tokens=5, max_retries=1) assert mock_client.responses.create.call_count == 1 -def test_openai_empty_text_response_raises_error() -> None: +@patch("time.sleep", return_value=None) +def test_openai_empty_text_response_raises_error(mock_sleep: MagicMock) -> None: """Test RuntimeError is raised when OpenAI returns empty text.""" mock_client = MagicMock() mock_response = MagicMock() mock_choice = MagicMock() - - mock_choice.text = "" + mock_choice.text = " \n" mock_response.choices = [mock_choice] mock_client.completions.create.return_value = mock_response @@ -249,6 +242,50 @@ def test_openai_empty_text_response_raises_error() -> None: provider.answer(prompt="Test empty response", model="gpt-3.5-turbo-instruct") mock_client.completions.create.assert_called_once() + mock_sleep.assert_not_called() + + +@patch("time.sleep", return_value=None) +def test_generate_text_exponential_backoff(mock_sleep: MagicMock) -> None: + """Ensure text generation retries respect exponential backoff caps at 30s.""" + client = MagicMock() + client.completions.create.side_effect = Exception("Random API failure") + + provider = OpenAIProvider(client) + with pytest.raises(RuntimeError, match="OpenAI request failed"): + provider.answer("test", max_retries=6) + + expected_calls = [call(2), call(4), call(8), call(16), call(30)] + mock_sleep.assert_has_calls(expected_calls) + assert mock_sleep.call_count == 5 + + +@patch("time.sleep", return_value=None) +def test_generate_text_custom_delay(mock_sleep: MagicMock) -> None: + """Ensure custom delay overrides default exponential backoff.""" + client = MagicMock() + client.completions.create.side_effect = [ + Exception("Temporary failure"), + MagicMock(choices=[MagicMock(text="success")]), + ] + + provider = OpenAIProvider(client) + result = provider.answer("test", max_retries=3, delay=1.5) + + assert result == "success" + mock_sleep.assert_called_once_with(1.5) + + +def test_generate_zero_retries_raises_fallback() -> None: + """Hit the end-of-function fallback RuntimeError by supplying max_retries=0.""" + client = MagicMock() + provider = OpenAIProvider(client) + + with pytest.raises( + RuntimeError, + match=re.escape("OpenAI text generation failed after max retries."), + ): + provider.answer("prompt", max_retries=0) # ------------------------------------------------------------------------- @@ -257,19 +294,15 @@ def test_openai_empty_text_response_raises_error() -> None: @patch("xwhy.providers.openai.Image.open") -@patch("os.makedirs") -@patch("time.time", return_value=1234567.89) def test_generate_image_b64_json_success( - mock_time: MagicMock, - mock_makedirs: MagicMock, mock_image_open: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test successful image generation using b64_json response format.""" client = MagicMock() mock_img_obj = MagicMock() mock_img_obj.b64_json = "dGVzdA==" # Base64 for "test" - mock_response = MagicMock() - mock_response.data = [mock_img_obj] + mock_response = MagicMock(data=[mock_img_obj]) client.images.generate.return_value = mock_response mock_img_instance = MagicMock(spec=Image.Image) @@ -277,11 +310,12 @@ def test_generate_image_b64_json_success( provider = OpenAIProvider(client) success, path = provider.generate_image( - prompt="test prompt", output_dir="fake_dir", model_name="test-model" + prompt="test prompt", output_dir=str(tmp_path), model_name="test-model" ) assert success is True - assert "openai_generated_1234567890.png" in path + assert "openai_generated_" in path + assert path.endswith(".png") mock_img_instance.save.assert_called_once() client.images.generate.assert_called_once_with( model="test-model", @@ -292,19 +326,17 @@ def test_generate_image_b64_json_success( @patch("xwhy.providers.openai.requests.get") @patch("xwhy.providers.openai.Image.open") -@patch("os.makedirs") def test_generate_image_url_success( - mock_makedirs: MagicMock, mock_image_open: MagicMock, mock_requests_get: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test successful image generation using url response format.""" client = MagicMock() mock_img_obj = MagicMock() mock_img_obj.b64_json = None mock_img_obj.url = "http://example.com/image.png" - mock_response = MagicMock() - mock_response.data = [mock_img_obj] + mock_response = MagicMock(data=[mock_img_obj]) client.images.generate.return_value = mock_response mock_request_response = MagicMock() @@ -315,13 +347,14 @@ def test_generate_image_url_success( mock_image_open.return_value = mock_img_instance provider = OpenAIProvider(client) - success, _ = provider.generate_image( + success, path = provider.generate_image( prompt="test prompt", - output_dir="fake_dir", + output_dir=str(tmp_path), response_format="url", ) assert success is True + assert path.startswith(str(tmp_path)) mock_requests_get.assert_called_once_with( "http://example.com/image.png", timeout=30 ) @@ -343,19 +376,17 @@ def test_edit_image_raises_file_not_found_error() -> None: @patch("builtins.open") @patch("xwhy.providers.openai.Image.open") @patch("os.path.exists", return_value=True) -@patch("os.makedirs") def test_edit_image_standard_api( - mock_makedirs: MagicMock, mock_path_exists: MagicMock, mock_image_open: MagicMock, mock_open_func: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test image editing using the standard images.edit endpoint with None format.""" client = MagicMock() mock_img_obj = MagicMock() mock_img_obj.b64_json = "dGVzdA==" - mock_response = MagicMock() - mock_response.data = [mock_img_obj] + mock_response = MagicMock(data=[mock_img_obj]) client.images.edit.return_value = mock_response mock_img_instance = MagicMock(spec=Image.Image) @@ -365,13 +396,12 @@ def test_edit_image_standard_api( success, _ = provider.edit_image( prompt="edit prompt", image_path="test.png", - output_dir="fake_dir", + output_dir=str(tmp_path), response_format=None, ) assert success is True client.images.edit.assert_called_once() - kwargs = client.images.edit.call_args.kwargs assert "response_format" not in kwargs @@ -379,19 +409,17 @@ def test_edit_image_standard_api( @patch("xwhy.providers.openai.image_to_base64") @patch("xwhy.providers.openai.Image.open") @patch("os.path.exists", return_value=True) -@patch("os.makedirs") def test_edit_image_use_generate_endpoint_with_url( - mock_makedirs: MagicMock, mock_path_exists: MagicMock, mock_image_open: MagicMock, mock_image_to_base64: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test editing image by routing through generation API with image_url payload.""" client = MagicMock() mock_img_obj = MagicMock() mock_img_obj.b64_json = "dGVzdA==" - mock_response = MagicMock() - mock_response.data = [mock_img_obj] + mock_response = MagicMock(data=[mock_img_obj]) client.images.generate.return_value = mock_response mock_image_to_base64.return_value = "data:image/png;base64,fake" @@ -402,7 +430,7 @@ def test_edit_image_use_generate_endpoint_with_url( success, _ = provider.edit_image( prompt="edit prompt", image_path="test.png", - output_dir="fake_dir", + output_dir=str(tmp_path), use_generate_for_edit=True, use_image_url=True, ) @@ -417,19 +445,17 @@ def test_edit_image_use_generate_endpoint_with_url( @patch("xwhy.providers.openai.image_to_base64") @patch("xwhy.providers.openai.Image.open") @patch("os.path.exists", return_value=True) -@patch("os.makedirs") def test_edit_image_use_generate_endpoint_with_image_key( - mock_makedirs: MagicMock, mock_path_exists: MagicMock, mock_image_open: MagicMock, mock_image_to_base64: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test editing image by routing through generation API with default image key.""" client = MagicMock() mock_img_obj = MagicMock() mock_img_obj.b64_json = "dGVzdA==" - mock_response = MagicMock() - mock_response.data = [mock_img_obj] + mock_response = MagicMock(data=[mock_img_obj]) client.images.generate.return_value = mock_response mock_image_to_base64.return_value = "data:image/png;base64,fake" @@ -440,7 +466,7 @@ def test_edit_image_use_generate_endpoint_with_image_key( success, _ = provider.edit_image( prompt="edit prompt", image_path="test.png", - output_dir="fake_dir", + output_dir=str(tmp_path), use_generate_for_edit=True, use_image_url=False, ) @@ -451,56 +477,58 @@ def test_edit_image_use_generate_endpoint_with_image_key( assert kwargs["extra_body"]["image"] == "data:image/png;base64,fake" -@patch("os.makedirs") +@patch("time.sleep", return_value=None) def test_execute_image_request_empty_data( - mock_makedirs: MagicMock, + mock_sleep: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test handling of empty image data returned from provider.""" client = MagicMock() - mock_response = MagicMock() - mock_response.data = [] + mock_response = MagicMock(data=[]) client.images.generate.return_value = mock_response provider = OpenAIProvider(client) mock_fallback = MagicMock(spec=Image.Image) provider._create_placeholder_image = MagicMock(return_value=mock_fallback) # type: ignore[method-assign] - success, _ = provider.generate_image("prompt", "out_dir") + success, path = provider.generate_image("prompt", str(tmp_path), max_retries=2) assert success is False + assert path.startswith(str(tmp_path)) provider._create_placeholder_image.assert_called_once() mock_fallback.save.assert_called_once() + mock_sleep.assert_called_once() -@patch("os.makedirs") +@patch("time.sleep", return_value=None) def test_execute_image_request_invalid_data( - mock_makedirs: MagicMock, + mock_sleep: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test handling of image data with no valid format fields.""" client = MagicMock() mock_img_obj = MagicMock() mock_img_obj.b64_json = None mock_img_obj.url = None - mock_response = MagicMock() - mock_response.data = [mock_img_obj] + mock_response = MagicMock(data=[mock_img_obj]) client.images.generate.return_value = mock_response provider = OpenAIProvider(client) - - # Mock fallback to prevent writing a real file to a patched directory mock_fallback = MagicMock(spec=Image.Image) provider._create_placeholder_image = MagicMock(return_value=mock_fallback) # type: ignore[method-assign] - success, _ = provider.generate_image("prompt", "out_dir") + success, _ = provider.generate_image("prompt", str(tmp_path), max_retries=1) assert success is False provider._create_placeholder_image.assert_called_once() mock_fallback.save.assert_called_once() + mock_sleep.assert_not_called() -@patch("os.makedirs") +@patch("time.sleep", return_value=None) def test_execute_image_request_exception_and_fallback( - mock_makedirs: MagicMock, + mock_sleep: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test that API exceptions are caught and fallback image is generated.""" client = MagicMock() @@ -510,28 +538,27 @@ def test_execute_image_request_exception_and_fallback( mock_fallback = MagicMock(spec=Image.Image) provider._create_placeholder_image = MagicMock(return_value=mock_fallback) # type: ignore[method-assign] - success, _ = provider.generate_image("prompt", "out_dir") + success, _ = provider.generate_image("prompt", str(tmp_path), max_retries=2) assert success is False provider._create_placeholder_image.assert_called_once() mock_fallback.save.assert_called_once() + mock_sleep.assert_called_once() -@patch("os.makedirs") +@patch("time.sleep", return_value=None) def test_execute_image_request_exception_no_fallback( - mock_makedirs: MagicMock, + mock_sleep: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test that API exceptions are caught and gracefully fail if no fallback.""" client = MagicMock() client.images.generate.side_effect = Exception("API failure") provider = OpenAIProvider(client) - - # Simulate a failed fallback assignment by returning None - # instead of trying to delete the class attribute provider._create_placeholder_image = MagicMock(return_value=None) # type: ignore[method-assign] - success, _ = provider.generate_image("prompt", "out_dir") + success, _ = provider.generate_image("prompt", str(tmp_path), max_retries=1) assert success is False provider._create_placeholder_image.assert_called_once() @@ -540,19 +567,17 @@ def test_execute_image_request_exception_no_fallback( @patch("xwhy.providers.openai.image_to_base64") @patch("xwhy.providers.openai.Image.open") @patch("os.path.exists", return_value=True) -@patch("os.makedirs") def test_edit_image_extra_body_image_url_exists( - mock_makedirs: MagicMock, mock_path_exists: MagicMock, mock_image_open: MagicMock, mock_image_to_base64: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test editing image when image_url is already provided in extra_body.""" client = MagicMock() mock_img_obj = MagicMock() mock_img_obj.b64_json = "dGVzdA==" - mock_response = MagicMock() - mock_response.data = [mock_img_obj] + mock_response = MagicMock(data=[mock_img_obj]) client.images.generate.return_value = mock_response mock_image_to_base64.return_value = "data:image/png;base64,fake" @@ -560,11 +585,10 @@ def test_edit_image_extra_body_image_url_exists( mock_image_open.return_value = mock_img_instance provider = OpenAIProvider(client) - success, _ = provider.edit_image( prompt="edit prompt", image_path="test.png", - output_dir="fake_dir", + output_dir=str(tmp_path), use_generate_for_edit=True, use_image_url=True, extra_body={"image_url": "pre_existing_url"}, @@ -579,19 +603,17 @@ def test_edit_image_extra_body_image_url_exists( @patch("xwhy.providers.openai.image_to_base64") @patch("xwhy.providers.openai.Image.open") @patch("os.path.exists", return_value=True) -@patch("os.makedirs") def test_edit_image_extra_body_image_exists( - mock_makedirs: MagicMock, mock_path_exists: MagicMock, mock_image_open: MagicMock, mock_image_to_base64: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test editing image when image key is already provided in extra_body.""" client = MagicMock() mock_img_obj = MagicMock() mock_img_obj.b64_json = "dGVzdA==" - mock_response = MagicMock() - mock_response.data = [mock_img_obj] + mock_response = MagicMock(data=[mock_img_obj]) client.images.generate.return_value = mock_response mock_image_to_base64.return_value = "data:image/png;base64,fake" @@ -599,11 +621,10 @@ def test_edit_image_extra_body_image_exists( mock_image_open.return_value = mock_img_instance provider = OpenAIProvider(client) - success, _ = provider.edit_image( prompt="edit prompt", image_path="test.png", - output_dir="fake_dir", + output_dir=str(tmp_path), use_generate_for_edit=True, use_image_url=False, extra_body={"image": "pre_existing_image"}, @@ -615,9 +636,10 @@ def test_edit_image_extra_body_image_exists( assert kwargs["extra_body"]["image"] == "pre_existing_image" -@patch("os.makedirs") +@patch("time.sleep", return_value=None) def test_execute_image_request_no_placeholder_attr( - mock_makedirs: MagicMock, + mock_sleep: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test handling failure when provider lacks placeholder generator attribute.""" client = MagicMock() @@ -625,32 +647,132 @@ def test_execute_image_request_no_placeholder_attr( provider = OpenAIProvider(client) - # Simulate hasattr(self, "_create_placeholder_image") evaluating to False with patch.object( OpenAIProvider, "_create_placeholder_image", new_callable=PropertyMock, ) as mock_prop: mock_prop.side_effect = AttributeError("Does not exist") - success, _ = provider.generate_image("prompt", "out_dir") + success, _ = provider.generate_image("prompt", str(tmp_path), max_retries=1) assert success is False -@patch("os.makedirs") +@patch("time.sleep", return_value=None) def test_execute_image_request_placeholder_not_an_image( - mock_makedirs: MagicMock, + mock_sleep: MagicMock, + tmp_path: Any, # noqa: ANN401 ) -> None: """Test failure fallback when placeholder generator returns non-Image.""" client = MagicMock() client.images.generate.side_effect = Exception("API failure") provider = OpenAIProvider(client) - - # Simulate the fallback function returning a string instead of PIL.Image.Image provider._create_placeholder_image = MagicMock(return_value="not_an_image") # type: ignore[method-assign] - success, _ = provider.generate_image("prompt", "out_dir") + success, _ = provider.generate_image("prompt", str(tmp_path), max_retries=1) + + assert success is False + provider._create_placeholder_image.assert_called_once() + + +@patch("time.sleep", return_value=None) +@patch("xwhy.providers.openai.Image.open") +def test_execute_image_request_retry_success( + mock_image_open: MagicMock, + mock_sleep: MagicMock, + tmp_path: Any, # noqa: ANN401 +) -> None: + """Test image generation successfully recovers after a failure with a delay.""" + client = MagicMock() + mock_img_obj = MagicMock() + mock_img_obj.b64_json = "dGVzdA==" + mock_response = MagicMock(data=[mock_img_obj]) + + client.images.generate.side_effect = [ + Exception("Temporary API glitch"), + mock_response, + ] + + mock_img_instance = MagicMock(spec=Image.Image) + mock_image_open.return_value = mock_img_instance + + provider = OpenAIProvider(client) + success, _ = provider.generate_image( + prompt="retry test", + output_dir=str(tmp_path), + max_retries=3, + delay=2.5, + ) + + assert success is True + assert client.images.generate.call_count == 2 + mock_sleep.assert_called_once_with(2.5) + + +@patch("time.sleep", return_value=None) +def test_execute_image_request_exponential_backoff( + mock_sleep: MagicMock, + tmp_path: Any, # noqa: ANN401 +) -> None: + """Ensure image generation retries correctly scale exponential sleep delays.""" + client = MagicMock() + client.images.generate.side_effect = Exception("Fatal Error") + + provider = OpenAIProvider(client) + provider._create_placeholder_image = MagicMock(return_value=None) # type: ignore[method-assign] + + success, _ = provider.generate_image("prompt", str(tmp_path), max_retries=5) assert success is False + expected_calls = [call(2), call(4), call(8), call(16)] + mock_sleep.assert_has_calls(expected_calls) + assert mock_sleep.call_count == 4 + + +@patch("time.sleep", return_value=None) +@patch("xwhy.providers.openai.Image.open") +def test_execute_image_request_image_open_returns_none( + mock_image_open: MagicMock, + mock_sleep: MagicMock, + tmp_path: Any, # noqa: ANN401 +) -> None: + """Test retry continuation when Image.open yields None. + + Verifies that if Image.open returns None on valid data, the loop + continues rather than breaking prematurely. + + Args: + mock_image_open: Mock for PIL.Image.open. + mock_sleep: Mock for time.sleep to avoid real delays. + tmp_path: Pytest temporary directory fixture. + + """ + client = MagicMock() + mock_img_obj = MagicMock() + mock_img_obj.b64_json = "dGVzdA==" + mock_response = MagicMock(data=[mock_img_obj]) + client.images.generate.return_value = mock_response + + # Force generated_img to remain None after a "successful" response. + mock_image_open.return_value = None + + provider = OpenAIProvider(client) + mock_fallback = MagicMock(spec=Image.Image) + provider._create_placeholder_image = MagicMock( # type: ignore[method-assign] + return_value=mock_fallback, + ) + + success, path = provider.generate_image( + prompt="prompt", + output_dir=str(tmp_path), + max_retries=2, + ) + + assert success is False + assert path.startswith(str(tmp_path)) + # Both attempts ran (no break occurred). + assert client.images.generate.call_count == 2 + mock_sleep.assert_called_once() provider._create_placeholder_image.assert_called_once() + mock_fallback.save.assert_called_once()