diff --git a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py index 99bfccb1a00..d3a5ade1cef 100644 --- a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py +++ b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py @@ -432,9 +432,12 @@ def test_get_request_body_nova_canvas_inference_profile_arn(): # Since we can't mock the actual model lookup, we'll test a simpler nova model instead # that we know the current logic can handle nova_model = "us.amazon.nova-canvas-v1:0" + + # Get the provider using the method from the handler + bedrock_provider = handler.get_bedrock_invoke_provider(model=nova_model) result = handler._get_request_body( - model=nova_model, bedrock_provider=None, prompt=prompt, optional_params=optional_params + model=nova_model, bedrock_provider=bedrock_provider, prompt=prompt, optional_params=optional_params ) assert result["taskType"] == "TEXT_IMAGE" @@ -484,10 +487,13 @@ def test_get_request_body_cross_region_inference_profile(): optional_params = {} # Cross-region inference profile format model = "us.amazon.nova-canvas-v1:0" + + # Get the provider using the method from the handler + bedrock_provider = handler.get_bedrock_invoke_provider(model=model) # This should work after the fix - cross-region format should be detected as 'nova' result = handler._get_request_body( - model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params + model=model, bedrock_provider=bedrock_provider, prompt=prompt, optional_params=optional_params ) assert result["taskType"] == "TEXT_IMAGE"