barunsaha commited on
Commit
dd9fc88
·
1 Parent(s): 41b5174

fix: propagate Azure OpenAI configs to the SlideDeckAI class and update tests

Browse files
.vscode/settings.json ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ {
2
+ "snyk.advanced.organization": "457f0b78-1951-4a11-b535-632990ba967e",
3
+ "snyk.advanced.autoSelectOrganization": true
4
+ }
app.py CHANGED
@@ -100,8 +100,8 @@ def are_all_inputs_valid(
100
  provider: str,
101
  selected_model: str,
102
  user_key: str,
103
- azure_deployment_url: str = '',
104
- azure_endpoint_name: str = '',
105
  azure_api_version: str = '',
106
  ) -> bool:
107
  """Validate user input and LLM selection.
@@ -111,8 +111,8 @@ def are_all_inputs_valid(
111
  provider: The LLM provider.
112
  selected_model: Name of the model.
113
  user_key: User-provided API key.
114
- azure_deployment_url: Azure OpenAI deployment URL.
115
- azure_endpoint_name: Azure OpenAI model endpoint.
116
  azure_api_version: Azure OpenAI API version.
117
 
118
  Returns:
@@ -135,8 +135,8 @@ def are_all_inputs_valid(
135
  provider,
136
  selected_model,
137
  user_key,
138
- azure_endpoint_name,
139
- azure_deployment_url,
140
  azure_api_version,
141
  ):
142
  handle_error(
@@ -277,15 +277,6 @@ with st.sidebar:
277
  disabled=bool(default_api_key),
278
  )
279
 
280
- # If a model was updated in the sidebar, make sure to update it in the SlideDeckAI instance
281
- if SLIDE_GENERATOR in st.session_state and llm_provider_to_use:
282
- try:
283
- st.session_state[SLIDE_GENERATOR].set_model(llm_provider_to_use, api_key_token)
284
- except Exception as e:
285
- logger.error('Failed to update model on existing SlideDeckAI: %s', e)
286
- # If updating fails, drop the stored instance so a new one is created
287
- st.session_state.pop(SLIDE_GENERATOR, None)
288
-
289
  # Additional configs for Azure OpenAI
290
  with st.expander('**Azure OpenAI-specific configurations**'):
291
  azure_endpoint = st.text_input(
@@ -307,6 +298,21 @@ with st.sidebar:
307
  value='2024-05-01-preview',
308
  )
309
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
310
  # Make slider with initial values
311
  page_range_slider = st.slider(
312
  'Specify a page range for the uploaded PDF file (if any):',
@@ -419,6 +425,9 @@ def set_up_chat_ui():
419
  st.session_state.get('start_page'),
420
  st.session_state.get('end_page'),
421
  ),
 
 
 
422
  )
423
  st.session_state[SLIDE_GENERATOR] = slide_generator
424
 
 
100
  provider: str,
101
  selected_model: str,
102
  user_key: str,
103
+ azure_endpoint_url: str = '',
104
+ azure_deployment_name: str = '',
105
  azure_api_version: str = '',
106
  ) -> bool:
107
  """Validate user input and LLM selection.
 
111
  provider: The LLM provider.
112
  selected_model: Name of the model.
113
  user_key: User-provided API key.
114
+ azure_endpoint_url: Azure OpenAI endpoint URL.
115
+ azure_deployment_name: Azure OpenAI deployment name.
116
  azure_api_version: Azure OpenAI API version.
117
 
118
  Returns:
 
135
  provider,
136
  selected_model,
137
  user_key,
138
+ azure_endpoint_url,
139
+ azure_deployment_name,
140
  azure_api_version,
141
  ):
142
  handle_error(
 
277
  disabled=bool(default_api_key),
278
  )
279
 
 
 
 
 
 
 
 
 
 
280
  # Additional configs for Azure OpenAI
281
  with st.expander('**Azure OpenAI-specific configurations**'):
282
  azure_endpoint = st.text_input(
 
298
  value='2024-05-01-preview',
299
  )
300
 
301
+ # If a model was updated in the sidebar, make sure to update it in the SlideDeckAI instance
302
+ if SLIDE_GENERATOR in st.session_state and llm_provider_to_use:
303
+ try:
304
+ st.session_state[SLIDE_GENERATOR].set_model(
305
+ llm_provider_to_use,
306
+ api_key=api_key_token,
307
+ azure_endpoint_url=azure_endpoint,
308
+ azure_deployment_name=azure_deployment,
309
+ azure_api_version=api_version,
310
+ )
311
+ except Exception as e:
312
+ logger.error('Failed to update model on existing SlideDeckAI: %s', e)
313
+ # If updating fails, drop the stored instance so a new one is created
314
+ st.session_state.pop(SLIDE_GENERATOR, None)
315
+
316
  # Make slider with initial values
317
  page_range_slider = st.slider(
318
  'Specify a page range for the uploaded PDF file (if any):',
 
425
  st.session_state.get('start_page'),
426
  st.session_state.get('end_page'),
427
  ),
428
+ azure_endpoint_url=azure_endpoint,
429
+ azure_deployment_name=azure_deployment,
430
+ azure_api_version=api_version,
431
  )
432
  st.session_state[SLIDE_GENERATOR] = slide_generator
433
 
src/slidedeckai/core.py CHANGED
@@ -78,6 +78,9 @@ class SlideDeckAI:
78
  pdf_path_or_stream=None,
79
  pdf_page_range=None,
80
  template_idx: int = 0,
 
 
 
81
  ):
82
  """Initialize the SlideDeckAI object.
83
 
@@ -88,6 +91,9 @@ class SlideDeckAI:
88
  pdf_path_or_stream: The path to a PDF file or a file-like object.
89
  pdf_page_range: A tuple representing the page range to use from the PDF file.
90
  template_idx: The index of the PowerPoint template to use.
 
 
 
91
 
92
  Raises:
93
  ValueError: If the model name is not in VALID_MODELS.
@@ -105,6 +111,9 @@ class SlideDeckAI:
105
  # Validate template_idx is within valid range
106
  num_templates = len(GlobalConfig.PPTX_TEMPLATE_FILES)
107
  self.template_idx: int = template_idx if 0 <= template_idx < num_templates else 0
 
 
 
108
  self.chat_history = ChatMessageHistory()
109
  self.last_response = None
110
  logger.info('Using model: %s', model)
@@ -124,6 +133,9 @@ class SlideDeckAI:
124
  model=llm_name,
125
  max_new_tokens=gcfg.get_max_output_tokens(self.model),
126
  api_key=self.api_key,
 
 
 
127
  )
128
 
129
  def _get_prompt_template(self, is_refinement: bool) -> str:
@@ -256,12 +268,22 @@ class SlideDeckAI:
256
 
257
  return path
258
 
259
- def set_model(self, model_name: str, api_key: str | None = None):
260
- """Set the LLM model (and API key) to use.
 
 
 
 
 
 
 
261
 
262
  Args:
263
  model_name: The name of the model to use.
264
  api_key: The API key for the LLM provider.
 
 
 
265
 
266
  Raises:
267
  ValueError: If the model name is not in VALID_MODELS.
@@ -271,8 +293,14 @@ class SlideDeckAI:
271
  f'Invalid model name: {model_name}. Must be one of: {", ".join(VALID_MODEL_NAMES)}.'
272
  )
273
  self.model = model_name
274
- if api_key:
275
  self.api_key = api_key
 
 
 
 
 
 
276
  logger.debug('Model set to: %s', model_name)
277
 
278
  def set_template(self, idx):
 
78
  pdf_path_or_stream=None,
79
  pdf_page_range=None,
80
  template_idx: int = 0,
81
+ azure_endpoint_url: str = '',
82
+ azure_deployment_name: str = '',
83
+ azure_api_version: str = '',
84
  ):
85
  """Initialize the SlideDeckAI object.
86
 
 
91
  pdf_path_or_stream: The path to a PDF file or a file-like object.
92
  pdf_page_range: A tuple representing the page range to use from the PDF file.
93
  template_idx: The index of the PowerPoint template to use.
94
+ azure_endpoint_url: Azure OpenAI endpoint URL (required when using Azure provider).
95
+ azure_deployment_name: Azure OpenAI deployment name (required when using Azure provider).
96
+ azure_api_version: Azure OpenAI API version (required when using Azure provider).
97
 
98
  Raises:
99
  ValueError: If the model name is not in VALID_MODELS.
 
111
  # Validate template_idx is within valid range
112
  num_templates = len(GlobalConfig.PPTX_TEMPLATE_FILES)
113
  self.template_idx: int = template_idx if 0 <= template_idx < num_templates else 0
114
+ self.azure_endpoint_url: str = azure_endpoint_url
115
+ self.azure_deployment_name: str = azure_deployment_name
116
+ self.azure_api_version: str = azure_api_version
117
  self.chat_history = ChatMessageHistory()
118
  self.last_response = None
119
  logger.info('Using model: %s', model)
 
133
  model=llm_name,
134
  max_new_tokens=gcfg.get_max_output_tokens(self.model),
135
  api_key=self.api_key,
136
+ azure_endpoint_url=self.azure_endpoint_url,
137
+ azure_deployment_name=self.azure_deployment_name,
138
+ azure_api_version=self.azure_api_version,
139
  )
140
 
141
  def _get_prompt_template(self, is_refinement: bool) -> str:
 
268
 
269
  return path
270
 
271
+ def set_model(
272
+ self,
273
+ model_name: str,
274
+ api_key: str | None = None,
275
+ azure_endpoint_url: str | None = None,
276
+ azure_deployment_name: str | None = None,
277
+ azure_api_version: str | None = None,
278
+ ):
279
+ """Set the LLM model (and optionally API key / Azure credentials) to use.
280
 
281
  Args:
282
  model_name: The name of the model to use.
283
  api_key: The API key for the LLM provider.
284
+ azure_endpoint_url: Azure OpenAI endpoint URL.
285
+ azure_deployment_name: Azure OpenAI deployment name.
286
+ azure_api_version: Azure OpenAI API version.
287
 
288
  Raises:
289
  ValueError: If the model name is not in VALID_MODELS.
 
293
  f'Invalid model name: {model_name}. Must be one of: {", ".join(VALID_MODEL_NAMES)}.'
294
  )
295
  self.model = model_name
296
+ if api_key is not None:
297
  self.api_key = api_key
298
+ if azure_endpoint_url is not None:
299
+ self.azure_endpoint_url = azure_endpoint_url
300
+ if azure_deployment_name is not None:
301
+ self.azure_deployment_name = azure_deployment_name
302
+ if azure_api_version is not None:
303
+ self.azure_api_version = azure_api_version
304
  logger.debug('Model set to: %s', model_name)
305
 
306
  def set_template(self, idx):
tests/unit/test_core.py CHANGED
@@ -18,6 +18,8 @@ from .test_utils import (
18
  with patch('transformers.BertTokenizer', patch_bert_tokenizer()):
19
  from slidedeckai.core import SlideDeckAI, _process_llm_chunk, _stream_llm_response
20
 
 
 
21
 
22
  @pytest.fixture
23
  def mock_env():
@@ -312,3 +314,110 @@ def test_topic_reset(slide_deck_ai):
312
  """Test that topic is retained after reset."""
313
  slide_deck_ai.reset()
314
  assert slide_deck_ai.topic == ''
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
18
  with patch('transformers.BertTokenizer', patch_bert_tokenizer()):
19
  from slidedeckai.core import SlideDeckAI, _process_llm_chunk, _stream_llm_response
20
 
21
+ from slidedeckai.global_config import GlobalConfig
22
+
23
 
24
  @pytest.fixture
25
  def mock_env():
 
314
  """Test that topic is retained after reset."""
315
  slide_deck_ai.reset()
316
  assert slide_deck_ai.topic == ''
317
+
318
+
319
+ def test_slide_deck_ai_init_azure_credentials():
320
+ """Test that SlideDeckAI stores Azure credentials provided at init time.
321
+
322
+ Regression test: previously, SlideDeckAI had no Azure params, so credentials
323
+ collected in the UI sidebar were silently dropped and never reached LiteLLM.
324
+ """
325
+ sda = SlideDeckAI(
326
+ model='[az]azure/open-ai',
327
+ topic='AI',
328
+ api_key='valid-key-12345',
329
+ azure_endpoint_url='https://test.openai.azure.com/',
330
+ azure_deployment_name='my-deployment',
331
+ azure_api_version='2024-05-01-preview',
332
+ )
333
+ assert sda.azure_endpoint_url == 'https://test.openai.azure.com/'
334
+ assert sda.azure_deployment_name == 'my-deployment'
335
+ assert sda.azure_api_version == '2024-05-01-preview'
336
+
337
+
338
+ @mock.patch('slidedeckai.core.llm_helper.get_provider_model')
339
+ @mock.patch('slidedeckai.core.llm_helper.get_litellm_llm')
340
+ def test_initialize_llm_azure_passes_credentials(mock_get_llm, mock_get_provider):
341
+ """Test that _initialize_llm() forwards Azure credentials to get_litellm_llm().
342
+
343
+ This is the core regression test: without the fix, _initialize_llm() called
344
+ get_litellm_llm() without azure_* params, causing a ValueError at stream time.
345
+ """
346
+ mock_get_provider.return_value = (GlobalConfig.PROVIDER_AZURE_OPENAI, 'azure/open-ai')
347
+ mock_get_llm.return_value = mock.Mock()
348
+
349
+ sda = SlideDeckAI(
350
+ model='[az]azure/open-ai',
351
+ topic='AI',
352
+ api_key='valid-key-12345',
353
+ azure_endpoint_url='https://test.openai.azure.com/',
354
+ azure_deployment_name='my-deployment',
355
+ azure_api_version='2024-05-01-preview',
356
+ )
357
+ sda._initialize_llm()
358
+
359
+ mock_get_llm.assert_called_once_with(
360
+ provider=GlobalConfig.PROVIDER_AZURE_OPENAI,
361
+ model='azure/open-ai',
362
+ max_new_tokens=mock.ANY,
363
+ api_key='valid-key-12345',
364
+ azure_endpoint_url='https://test.openai.azure.com/',
365
+ azure_deployment_name='my-deployment',
366
+ azure_api_version='2024-05-01-preview',
367
+ )
368
+
369
+
370
+ @mock.patch.dict(
371
+ 'slidedeckai.core.GlobalConfig.VALID_MODELS',
372
+ {
373
+ '[az]azure/open-ai': {'description': 'azure', 'max_new_tokens': 8192, 'paid': True},
374
+ },
375
+ )
376
+ def test_set_model_updates_azure_credentials():
377
+ """Test that set_model() updates Azure credentials when provided."""
378
+ sda = SlideDeckAI(
379
+ model='[az]azure/open-ai',
380
+ topic='AI',
381
+ api_key='old-key-12345',
382
+ azure_endpoint_url='https://old.openai.azure.com/',
383
+ azure_deployment_name='old-deployment',
384
+ azure_api_version='2024-02-01',
385
+ )
386
+
387
+ sda.set_model(
388
+ '[az]azure/open-ai',
389
+ api_key='new-key-12345',
390
+ azure_endpoint_url='https://new.openai.azure.com/',
391
+ azure_deployment_name='new-deployment',
392
+ azure_api_version='2024-05-01-preview',
393
+ )
394
+
395
+ assert sda.api_key == 'new-key-12345'
396
+ assert sda.azure_endpoint_url == 'https://new.openai.azure.com/'
397
+ assert sda.azure_deployment_name == 'new-deployment'
398
+ assert sda.azure_api_version == '2024-05-01-preview'
399
+
400
+
401
+ @mock.patch.dict(
402
+ 'slidedeckai.core.GlobalConfig.VALID_MODELS',
403
+ {
404
+ '[az]azure/open-ai': {'description': 'azure', 'max_new_tokens': 8192, 'paid': True},
405
+ },
406
+ )
407
+ def test_set_model_preserves_azure_credentials_when_not_provided():
408
+ """Test that set_model() keeps existing Azure credentials when None is passed."""
409
+ sda = SlideDeckAI(
410
+ model='[az]azure/open-ai',
411
+ topic='AI',
412
+ api_key='valid-key-12345',
413
+ azure_endpoint_url='https://test.openai.azure.com/',
414
+ azure_deployment_name='my-deployment',
415
+ azure_api_version='2024-05-01-preview',
416
+ )
417
+
418
+ # Call set_model without Azure params — existing values must be preserved
419
+ sda.set_model('[az]azure/open-ai')
420
+
421
+ assert sda.azure_endpoint_url == 'https://test.openai.azure.com/'
422
+ assert sda.azure_deployment_name == 'my-deployment'
423
+ assert sda.azure_api_version == '2024-05-01-preview'
tests/unit/test_llm_helper.py CHANGED
@@ -259,3 +259,60 @@ def test_stream_litellm_completion_message_format(mock_litellm):
259
 
260
  assert result == ['Alternative format']
261
  mock_litellm.completion.assert_called_once()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
259
 
260
  assert result == ['Alternative format']
261
  mock_litellm.completion.assert_called_once()
262
+
263
+
264
+ def test_stream_litellm_completion_azure_missing_deployment():
265
+ """Test that stream_litellm_completion raises ValueError when Azure deployment name is empty.
266
+
267
+ This is the precise condition that caused the runtime error when Azure OpenAI was selected
268
+ but credentials were not propagated through SlideDeckAI._initialize_llm().
269
+ """
270
+ messages = [{'role': 'user', 'content': 'Test'}]
271
+ with pytest.raises(ValueError, match='Azure deployment name is required'):
272
+ list(
273
+ stream_litellm_completion(
274
+ provider=GlobalConfig.PROVIDER_AZURE_OPENAI,
275
+ model='gpt-4',
276
+ messages=messages,
277
+ max_tokens=100,
278
+ api_key='valid-key-12345',
279
+ azure_endpoint_url='https://test.openai.azure.com/',
280
+ azure_deployment_name='', # Empty — the missing credential
281
+ azure_api_version='2024-05-01-preview',
282
+ )
283
+ )
284
+
285
+
286
+ @patch('slidedeckai.helpers.llm_helper.stream_litellm_completion')
287
+ def test_get_litellm_llm_azure_passes_credentials(mock_stream):
288
+ """Test that get_litellm_llm forwards Azure credentials to stream_litellm_completion.
289
+
290
+ Regression test: SlideDeckAI._initialize_llm() previously called get_litellm_llm()
291
+ without Azure params, so LiteLLMWrapper was created with empty strings and the
292
+ deployment-name check in stream_litellm_completion raised a ValueError at call time.
293
+ """
294
+ mock_stream.return_value = iter(['Azure response'])
295
+
296
+ llm = get_litellm_llm(
297
+ provider=GlobalConfig.PROVIDER_AZURE_OPENAI,
298
+ model='gpt-4',
299
+ max_new_tokens=100,
300
+ api_key='valid-key-12345',
301
+ azure_endpoint_url='https://test.openai.azure.com/',
302
+ azure_deployment_name='my-deployment',
303
+ azure_api_version='2024-05-01-preview',
304
+ )
305
+
306
+ result = list(llm.stream('Hello'))
307
+ assert result == ['Azure response']
308
+
309
+ mock_stream.assert_called_once_with(
310
+ provider=GlobalConfig.PROVIDER_AZURE_OPENAI,
311
+ model='gpt-4',
312
+ messages=[{'role': 'user', 'content': 'Hello'}],
313
+ max_tokens=100,
314
+ api_key='valid-key-12345',
315
+ azure_endpoint_url='https://test.openai.azure.com/',
316
+ azure_deployment_name='my-deployment',
317
+ azure_api_version='2024-05-01-preview',
318
+ )