Handles the main code generation logic
| 770 | |
| 771 | |
| 772 | class CodeGenerationMiddleware(Middleware): |
| 773 | """Handles the main code generation logic""" |
| 774 | |
| 775 | async def process( |
| 776 | self, context: PipelineContext, next_func: Callable[[], Awaitable[None]] |
| 777 | ) -> None: |
| 778 | try: |
| 779 | assert context.extracted_params is not None |
| 780 | |
| 781 | # Select models (handles video mode internally) |
| 782 | model_selector = ModelSelectionStage(context.throw_error) |
| 783 | context.variant_models = await model_selector.select_models( |
| 784 | generation_type=context.extracted_params.generation_type, |
| 785 | input_mode=context.extracted_params.input_mode, |
| 786 | openai_api_key=context.extracted_params.openai_api_key, |
| 787 | anthropic_api_key=context.extracted_params.anthropic_api_key, |
| 788 | gemini_api_key=context.extracted_params.gemini_api_key, |
| 789 | ) |
| 790 | if IS_DEBUG_ENABLED: |
| 791 | await context.send_message( |
| 792 | "variantModels", |
| 793 | None, |
| 794 | 0, |
| 795 | {"models": [model.value for model in context.variant_models]}, |
| 796 | None, |
| 797 | ) |
| 798 | |
| 799 | generation_stage = AgenticGenerationStage( |
| 800 | send_message=context.send_message, |
| 801 | openai_api_key=context.extracted_params.openai_api_key, |
| 802 | openai_base_url=context.extracted_params.openai_base_url, |
| 803 | anthropic_api_key=context.extracted_params.anthropic_api_key, |
| 804 | gemini_api_key=context.extracted_params.gemini_api_key, |
| 805 | replicate_api_key=context.extracted_params.replicate_api_key, |
| 806 | should_generate_images=context.extracted_params.should_generate_images, |
| 807 | file_state=context.extracted_params.file_state, |
| 808 | asset_base_url=context.extracted_params.asset_base_url, |
| 809 | option_codes=context.extracted_params.option_codes, |
| 810 | ) |
| 811 | |
| 812 | context.variant_completions = await generation_stage.process_variants( |
| 813 | variant_models=context.variant_models, |
| 814 | prompt_messages=context.prompt_messages, |
| 815 | ) |
| 816 | |
| 817 | # Check if all variants failed |
| 818 | if len(context.variant_completions) == 0: |
| 819 | await context.throw_error( |
| 820 | "Error generating code. Please contact support." |
| 821 | ) |
| 822 | return # Don't continue the pipeline |
| 823 | |
| 824 | # Convert to list format |
| 825 | context.completions = [] |
| 826 | for i in range(len(context.variant_models)): |
| 827 | if i in context.variant_completions: |
| 828 | context.completions.append(context.variant_completions[i]) |
| 829 | else: |