chore: import upstream snapshot with attribution
This commit is contained in:
@@ -0,0 +1,201 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
|
||||
Copyright [yyyy] [name of copyright owner]
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
@@ -0,0 +1,21 @@
|
||||
# Sample Applications
|
||||
|
||||
This directory contains sample applications that demonstrate how to use the Gemini API in Vertex AI.
|
||||
|
||||
## Applications
|
||||
|
||||
- [Accelerating Product Innovation](accelerating_product_innovation/): A solution for product category and brand owners to accelerate product innovation using Generative AI.
|
||||
- [End-to-End Gen AI App Starter Pack](e2e-gen-ai-app-starter-pack/): This project has moved to https://github.com/GoogleCloudPlatform/agent-starter-pack.
|
||||
- [Finance Advisor with Spanner](finance-advisor-spanner/): A demo application that showcases how to use Spanner and Vertex AI to build a financial advisor application.
|
||||
- [FixMyCar](fixmycar/): A retrieval-augmented generation (RAG) sample application to troubleshoot your car using the owner's manual.
|
||||
- [Gemini Hallucination Check](gemini-hallcheck/): A confidence-targeted, abstention-aware hallucination evaluator for Gemini.
|
||||
- [Gemini with Mesop on Cloud Run](gemini-mesop-cloudrun/): A sample application that demonstrates how to use the Mesop UI framework with the Gemini API on Cloud Run.
|
||||
- [Gemini with Quart on Cloud Run](gemini-quart-cloudrun/): A sample application that demonstrates non-blocking communication with Quart and the Gemini Live API on Cloud Run.
|
||||
- [Gemini with Streamlit on Cloud Run](gemini-streamlit-cloudrun/): A sample application that demonstrates how to use the Streamlit framework with the Gemini API on Cloud Run.
|
||||
- [GenWealth](genwealth/): A demo application for a fictional financial services company that showcases how to build trustworthy Gen AI features into existing applications using AlloyDB AI, Vertex AI, Cloud Run, and Cloud Functions.
|
||||
- [Image Bash JAM](image-bash-jam/): A collection of bash scripts to test Gemini from the command line, including image and audio examples.
|
||||
- [LlamaDeploy on Cloud Run](llamadeploy-on-cloud-run/): A LlamaIndex Workflow application that demonstrates how to deploy and interact with Llama workflows using the `llama-deploy` library and deploying the service on Cloud Run.
|
||||
- [LlamaIndex RAG](llamaindex-rag/): An advanced Retrieval-Augmented Generation (RAG) system using LlamaIndex and Google Cloud Vertex AI for rapid prototyping and experimentation.
|
||||
- [Photo Discovery](photo-discovery/): A demo that integrates a Vertex AI Agent with a multi-platform Flutter app to identify landmarks and merchandise from photos.
|
||||
- [Quickbot](quickbot/): An innovative, out-of-the-box solution enabling users to deploy sophisticated AI Agents as full-stack cloud applications on their own Google Cloud Platform (GCP) accounts, entirely without requiring any coding expertise.
|
||||
- [SWOT Agent](swot-agent/): A web application that performs automated SWOT analysis (Strengths, Weaknesses, Opportunities, Threats) using the Gemini 2.0 Flash model and the Pydantic AI agent framework.
|
||||
@@ -0,0 +1,95 @@
|
||||
# Self-paced environment setup
|
||||
|
||||
1. Sign-in to the [Google Cloud Console](http://console.cloud.google.com/) and create a new project or reuse an existing one. If you don't already have a Gmail or Google Workspace account, you must [create one](https://accounts.google.com/SignUp).
|
||||
|
||||

|
||||
|
||||
- The **Project name** is the display name for this project's participants. It is a character string not used by Google APIs. You can always update it.
|
||||
|
||||
- The **Project ID** is unique across all Google Cloud projects and is immutable (cannot be changed after it has been set). The Cloud Console auto-generates a unique string; usually you don't care what it is. In most codelabs, you'll need to reference your `Project ID` (typically identified as PROJECT_ID). If you don't like the generated ID, you might generate another random one. Alternatively, you can try your own, and see if it's available. It can't be changed after this step and remains for the duration of the project.
|
||||
|
||||
- For your information, there is a third value, a **Project Number**, which some APIs use. Learn more about all three of these values in the documentation.
|
||||
|
||||
2. Next, you'll need to [enable billing](https://console.cloud.google.com/billing) in the Cloud Console to use Cloud resources/APIs. Running through this codelab won't cost much, if anything at all. To shut down resources to avoid incurring billing beyond this tutorial, you can delete the resources you created or delete the project. New Google Cloud users are eligible for the [$300 USD Free Trial](http://cloud.google.com/free) program.
|
||||
|
||||
## Start Cloud Shell
|
||||
|
||||
While Google Cloud can be operated remotely from your laptop, in this codelab you will be using [Google Cloud Shell](https://cloud.google.com/cloud-shell/), a command line environment running in the Cloud.
|
||||
|
||||
From the [Google Cloud Console](https://console.cloud.google.com/), click the Cloud Shell icon on the top right toolbar:
|
||||
|
||||

|
||||
|
||||
It should only take a few moments to provision and connect to the environment. When it is finished, you should see something like this:
|
||||
|
||||

|
||||
|
||||
This virtual machine is loaded with all the development tools you'll need. It offers a persistent 5GB home directory, and runs on Google Cloud, greatly enhancing network performance and authentication. All of your work in this codelab can be done within a browser. You do not need to install anything.
|
||||
|
||||
Once connected to Cloud Shell, you should see that you are already authenticated and that the project is already set to your project ID.
|
||||
|
||||
Run the following command in Cloud Shell to confirm that you are authenticated:
|
||||
|
||||
Once connected to Cloud Shell, you should see that you are already authenticated and that the project is already set to your `PROJECT_ID``.
|
||||
|
||||
```bash
|
||||
gcloud auth list
|
||||
```
|
||||
|
||||
Command output:
|
||||
|
||||
```bash
|
||||
Credentialed accounts:
|
||||
- <myaccount>@<mydomain>.com (active)
|
||||
```
|
||||
|
||||
```bash
|
||||
gcloud config list project
|
||||
```
|
||||
|
||||
Command output:
|
||||
|
||||
```bash
|
||||
[core]
|
||||
project = <PROJECT_ID>
|
||||
```
|
||||
|
||||
If, for some reason, the project is not set, simply issue the following command:
|
||||
|
||||
```bash
|
||||
gcloud config set project <PROJECT_ID>
|
||||
```
|
||||
|
||||
Cloud Shell also sets some environment variables by default, which may be useful as you run future commands.
|
||||
|
||||
```bash
|
||||
echo $GOOGLE_CLOUD_PROJECT
|
||||
```
|
||||
|
||||
Command output:
|
||||
|
||||
```bash
|
||||
<PROJECT_ID>
|
||||
```
|
||||
|
||||
## Enable the Google Cloud APIs
|
||||
|
||||
In order to use the various services we will need throughout this project, we will enable a few APIs. We will do so by launching the following command in Cloud Shell:
|
||||
|
||||
```bash
|
||||
gcloud services enable cloudbuild.googleapis.com cloudfunctions.googleapis.com run.googleapis.com logging.googleapis.com storage-component.googleapis.com aiplatform.googleapis.com
|
||||
```
|
||||
|
||||
After some time, you should see the operation finish successfully:
|
||||
|
||||
```bash
|
||||
Operation "operations/acf.5c5ef4f6-f734-455d-b2f0-ee70b5a17322" finished successfully.
|
||||
```
|
||||
|
||||
## Clone the Repository
|
||||
|
||||
We've put all the samples you need for this project into a Git repo in the `sample-apps` folder. Clone the repo in Cloud Shell using the following command:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/GoogleCloudPlatform/generative-ai.git
|
||||
```
|
||||
@@ -0,0 +1,27 @@
|
||||
# Copyright 2023 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# https://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
|
||||
[server]
|
||||
runOnSave = true
|
||||
|
||||
port = 8080
|
||||
|
||||
[browser]
|
||||
gatherUsageStats = false
|
||||
serverAddress = "0.0.0.0"
|
||||
|
||||
[theme]
|
||||
base="light"
|
||||
primaryColor="#0066cc"
|
||||
@@ -0,0 +1,20 @@
|
||||
|
||||
[[pages]]
|
||||
path = "app/home_page.py"
|
||||
name = "Home Page"
|
||||
|
||||
[[pages]]
|
||||
path = "app/pages/resources.py"
|
||||
name = "Resources"
|
||||
|
||||
[[pages]]
|
||||
path = "app/pages/product_insights.py"
|
||||
name = "Product Insights"
|
||||
|
||||
[[pages]]
|
||||
path = "app/pages/product_generation.py"
|
||||
name = "Product Generation"
|
||||
|
||||
[[pages]]
|
||||
path = "app/pages/edit_image.py"
|
||||
name = "Edit Image"
|
||||
@@ -0,0 +1,138 @@
|
||||
# Generative AI Demos - Accelerating Product Innovation
|
||||
|
||||
## Introduction
|
||||
|
||||
This solution is for product category and brand owners, product R&D analysts, marketers, and any personas owning the development of new products in Retail, or any other vertical, where there is a need for a rapid and adaptive pace of new product and new product variant innovation. The solution enables users to leverage Generative AI's purely creative capabilities to new ideas and new concepts for new products.
|
||||
|
||||
In fast and dynamic markets there is a need to,
|
||||
|
||||
- shorten lead time for new product development; and
|
||||
- achieve draft ideas in bulk with minimal or without human-in-the-loop
|
||||
|
||||
This Streamlit-based solution empowers product managers, R&D specialists, and marketers to harness the power of Generative AI for accelerated product development, discover how to rapidly generate new product concepts, address market trends, and ensure regulatory compliance within the retail sector and beyond.
|
||||
|
||||
## Getting Started
|
||||
|
||||
To access the application, follow the steps in `Setup.md`. Once you have the solution running, follow these procedures within the application to generate innovative ideas for any product:
|
||||
|
||||
- **Navigate to the 'Resources' page.**
|
||||
- **Project Setup**: Begin by creating a new project or selecting an existing one.
|
||||
- **Document Upload**: Add relevant research and data files in accepted formats.
|
||||
- **For Q&A on uploaded files, navigate to 'Product Insights' page**
|
||||
- **Insight Generation**: Ask questions to extract critical information from your data.
|
||||
- **For Product Idea generation, navigate to 'Product Generation' page**
|
||||
- **Concept Creation**: Initiate product generation using sample queries or create your own custom query.
|
||||
- **Refinement**: Select features, experiment with combinations, and regenerate results until the desired product concept is achieved.
|
||||
|
||||
## Application Workflow
|
||||
|
||||
1. **Project Setup**
|
||||
You have the option to either create a new project or select an existing project in the resources page of the application for generating product insights. Choose one of the following steps:
|
||||
|
||||
- **New Project Creation**:
|
||||
- Initiate a new project within the application by giving it a descriptive name (e.g., "2024 Sunscreen Innovation" or "Hair care Line Extension").
|
||||
- Once a project is created, the document upload process is mandatory before proceeding to analysis features.
|
||||
- **Existing Project Modification**:
|
||||
- Select an existing project from a list.
|
||||
- View previously uploaded documents associated with the project.
|
||||
- Use the following actions:
|
||||
- Add Files: Upload new documents relevant to the project.
|
||||
- Remove Files: Delete documents that are outdated or no longer relevant to the project goals.
|
||||
- Delete Project: Entirely remove the project and all the resources associated with it.
|
||||
|
||||
### 1. Document Upload
|
||||
|
||||
<p align="center">
|
||||
<img src="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/readme_images/resource_upload.gif" alt="Image Description" width="600"/>
|
||||
</p>
|
||||
|
||||
- **Objective**: Integrate your critical project data for analysis with the solution.
|
||||
- **Accepted File Formats**:
|
||||
Various types of documents such as market research reports, consumer feedback surveys, Internal trend analyses, regulatory guidelines can be uploaded. Accepted file formats include:
|
||||
- Documents: .pdf, .doc, .docx
|
||||
- Spreadsheets: .xlsx, .csv
|
||||
- Text Files: .txt, .md
|
||||
|
||||
### 2. Product Insights
|
||||
|
||||
- **Dynamic Suggestions**
|
||||
|
||||
<p align="center">
|
||||
<img src="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/readme_images/suggestions.png" alt="Image Description" width="600"/>
|
||||
</p>
|
||||
|
||||
- **Objective**: This feature generates contextually relevant suggestions designed to extract essential product insights and attributes companies deem important.
|
||||
- **Functionality**: The system leverages information embedded in the documents you've uploaded (market research, surveys, etc.) to tailor sample queries.
|
||||
- **Usage**: Choose to utilize a suggested sample query for immediate results. Customize the sample query or input your own questions to drive a more focused analysis.
|
||||
|
||||
- **Generative AI Powered Product Insights**
|
||||
- **Objective**: This feature unlocks insights hidden within your uploaded documents. You pose questions in natural language, and the system retrieves contextually relevant answers.
|
||||
- **Functionality**: Rapid insight extraction: Obtain actionable answers directly from your data without time-consuming manual analysis.
|
||||
- **Usage**: Type your question in natural language as if you were asking a domain expert.
|
||||
- **Question Types**:
|
||||
- Facts
|
||||
- Trends
|
||||
- Relationships
|
||||
- Comparisons
|
||||
|
||||
<p align="center">
|
||||
<img src="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/readme_images/insights.png" alt="Image Description" width="600"/>
|
||||
</p>
|
||||
|
||||
<p align="center">
|
||||
<img src="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/readme_images/follow_up_qs.png" alt="Image Description" width="600"/>
|
||||
</p>
|
||||
|
||||
- **Benefits**
|
||||
- **Document Compatibility**: The AI-Powered Question Answering functionality is optimized for document types including: Market Research Reports, Consumer Feedback Surveys, Internal Trend Analyses.
|
||||
|
||||
### 3. Product Generation
|
||||
|
||||
- **Dynamic Sample Queries**
|
||||
|
||||
<p align="center">
|
||||
<img src="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/readme_images/queries.png" alt="Image Description" width="600"/>
|
||||
</p>
|
||||
|
||||
- **Objective**: This feature generates contextually relevant questions designed to extract essential product attributes consumers deem important within your chosen category.
|
||||
- **Functionality**: The system leverages information embedded in the documents you've uploaded (market research, surveys, etc.) to tailor sample queries.
|
||||
- **Usage**: Choose to utilize a suggested sample query for immediate results. Customize the sample query or input your own questions to drive a more focused analysis.
|
||||
|
||||
- **Key Feature Extraction**
|
||||
|
||||
- **Objective**: Isolates the most frequently mentioned or prioritized product features that emerge from the analysis of your uploaded documents.
|
||||
- **Functionality**: Employs text analysis techniques to extract key features. Extracts ingredients, benefits, claims, sensory terms, usage occasions, and more.
|
||||
- **Output**: Creates a curated list of extracted features.
|
||||
|
||||
- **Feature Selection**
|
||||
|
||||
<p align="center">
|
||||
<img src="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/readme_images/features.png" alt="Image Description" width="600"/>
|
||||
</p>
|
||||
|
||||
- **Objective**: Gives you direct control in constructing your ideal product concept.
|
||||
- **Functionality**: Presents the extracted list of key features for selection. Enables both single feature selection and the creation of novel feature combinations.
|
||||
|
||||
- **Product Concept Generation**
|
||||
|
||||
- **Objective**: Transforms selected features into holistic product concepts, going beyond simply listing ingredients or attributes.
|
||||
- **Functionality**: Utilizes Generative AI, trained on product descriptions, marketing material, and relevant category data. Presents multiple plausible product ideas based on your selections.
|
||||
- **Output**: Presents a series of product ideas featuring the AI-generated concepts, helping to visualize the creative possibilities.
|
||||
|
||||
<p align="center">
|
||||
<img src="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/readme_images/product.png" alt="Image Description" width="600"/>
|
||||
</p>
|
||||
|
||||
### 4. Selective Regeneration
|
||||
|
||||
- **Objective**: Enables users to zero in on specific areas of the product concept that they wish to modify without re-generating the whole concept from scratch.
|
||||
- **Functionality**: Text Modification, Image Modification, Whole Product Regeneration.
|
||||
|
||||
<p align="center">
|
||||
<img src="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/readme_images/text_regen.gif" alt="Image Description" width="600"/>
|
||||
</p>
|
||||
|
||||
### 5. Export Content
|
||||
|
||||
- **Objective**: Once you've refined your product concept through the iterative process, the solution provides seamless ways to share your work and integrate it into your broader product development workflows.
|
||||
- **Functionalities**: PDF Export, Export/Download, Email Export.
|
||||
@@ -0,0 +1,84 @@
|
||||
# Setup Steps
|
||||
|
||||
The solution supports both local execution and deployment to Cloud Run environments.
|
||||
|
||||
## Running Locally
|
||||
|
||||
1. **Clone Repository**: Clone the repository containing the solution's code to your local machine.
|
||||
2. **Install Dependencies**: Navigate to the project directory and install dependencies from requirements.txt using pip.
|
||||
3. **Set Environment Variables**: Create a `.env` file in the project directory and populate it with necessary environment variables:
|
||||
|
||||
```plaintext
|
||||
PROJECT_ID=<Your_Google_Cloud_Project_ID>
|
||||
LOCATION=<Desired_Location>
|
||||
REGION=<Desired_Google_Cloud_Project_Region>
|
||||
YOUR_EMAIL=<Your_Email_Address_Associated_With_Google_Cloud_Project>
|
||||
PROJECT_NUMBER=<Your_Google_Cloud_Project_Number>
|
||||
```
|
||||
|
||||
4. **Run the Application**: Execute the following command to run the application locally:
|
||||
|
||||
```bash
|
||||
python -m streamlit run app/Home.py
|
||||
```
|
||||
|
||||
5. **Access the Application**: Once the application is running locally, access it through a web browser using the specified local host address and port.
|
||||
|
||||
## Deployment Steps
|
||||
|
||||
Follow the below steps to deploy the solution to Cloud Run environment.
|
||||
|
||||
### Google Cloud Storage Setup
|
||||
|
||||
1. **Create Bucket**: Manually create GCS bucket 'product_innovation_bucket' using either the Google Cloud Console or command-line tools (gsutil). The bucket is necessary for:
|
||||
|
||||
- `document_uploads`: Stores market research, surveys, trend reports, etc.
|
||||
- `generated_products`: Stores output images, descriptions, etc.
|
||||
- `image_edits`: Stores intermediate/modified images during the regeneration process.
|
||||
|
||||
**Using the Google Cloud Console (Web Interface):**
|
||||
|
||||
- Navigate to the [Google Cloud Storage section](https://console.cloud.google.com/storage/browser) of your Google Cloud console.
|
||||
- Click on the "Create Bucket" button.
|
||||
- Provide the following details:
|
||||
- **Name**: Enter 'product_innovation_bucket' as bucket name.
|
||||
- **Location**: Choose a region closest to where your solution will operate for best performance.
|
||||
- **Storage Class**: Select the class based on frequency of access and cost considerations.
|
||||
- **Advanced Settings**: Adjust encryption, access control, etc., if necessary.
|
||||
- Click "Create".
|
||||
|
||||
**Using the 'gsutil' Command-Line Tool:**
|
||||
|
||||
- Ensure that you have the gcloud SDK installed and 'gsutil' configured.
|
||||
- Run the following command in your terminal, replacing `<region>` with your desired bucket location:
|
||||
|
||||
```bash
|
||||
gcloud storage buckets create --location <region> gs://product_innovation_bucket
|
||||
```
|
||||
|
||||
### Environment Setup
|
||||
|
||||
1. Ensure that the `configure_resources.sh` script is in the directory containing the solution's code.
|
||||
2. Create a `.env` file within the same directory and populate it with the same values as mentioned above.
|
||||
3. Create a `.env` file in the 'cloud_functions' directory and populate it with necessary environment variables as mentioned above.
|
||||
|
||||
### Execute the Script
|
||||
|
||||
- Open a terminal and navigate to the directory containing the script and `env.txt` file.
|
||||
- Run the script `configure-resources.sh` using the command:
|
||||
|
||||
```bash
|
||||
sh configure-resources.sh
|
||||
```
|
||||
|
||||
- The script will:
|
||||
- Parse `.env` to obtain project details.
|
||||
- Initialize gcloud and set project configuration.
|
||||
- Set up a service account with necessary IAM roles.
|
||||
- Deploy Cloud Functions (`imagen-call`, `gemini-call`, `text-embedding`).
|
||||
- Capture URLs for deployed Cloud Functions.
|
||||
- Deploy the main application to Cloud Run.
|
||||
- Ensure that the service account has been created and manually grant the following roles to the created service account `retail-accelerating-prod-i-982@[PROJECT_ID].iam.gserviceaccount.com`:
|
||||
- Service account user
|
||||
- Cloud Run Admin
|
||||
- Cloud Storage Admin
|
||||
@@ -0,0 +1,76 @@
|
||||
# -----------------------------------------------------------------------------
|
||||
# Global Configurations
|
||||
# -----------------------------------------------------------------------------
|
||||
[global]
|
||||
|
||||
bucket_name = "product_innovation_bucket"
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Translation options
|
||||
# -----------------------------------------------------------------------------
|
||||
[translate_api]
|
||||
|
||||
# Include or remove languages as needed. Use the this documentation as
|
||||
# reference: https://cloud.google.com/translate/docs/languages
|
||||
Spanish = "es"
|
||||
Chinese = "zh"
|
||||
German = "de"
|
||||
Japanese = "ja"
|
||||
Portuguese = "pt"
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Configurations for page: Home.py
|
||||
# -----------------------------------------------------------------------------
|
||||
[pages.home]
|
||||
|
||||
page_title = "Gen AI for Product Innovation"
|
||||
page_icon = "https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/favicon.png"
|
||||
sidebar_image_path = "https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/main_logo.png"
|
||||
home_img_1 = "https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/intro_home_1.png"
|
||||
home_img_2 = "https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/intro_home_2.png"
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Configurations for page: 1_Resources.py
|
||||
# -----------------------------------------------------------------------------
|
||||
[pages.1_Resources]
|
||||
|
||||
page_title = "Resources"
|
||||
page_icon = "https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/favicon.png"
|
||||
sidebar_image_path = "https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/main_logo.png"
|
||||
resources_img="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/resources_img.png"
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Configurations for page: 2_ Marketing_Insights.py
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
[pages.2_Marketing_Insights]
|
||||
page_title = "Marketing Insights"
|
||||
page_icon = "./images/favicon.png"
|
||||
sidebar_image_path = "https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/main_logo.png"
|
||||
prod_insights_1="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/prod_insights_1.png"
|
||||
prod_insights_2="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/prod_insights_2.png"
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Configurations for page: 3_Generations.py
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
[pages.3_Generations]
|
||||
page_title = "Generations"
|
||||
page_icon = "https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/favicon.png"
|
||||
sidebar_image_path = "https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/main_logo.png"
|
||||
prod_gen_img="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/prod_gen_img.png"
|
||||
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Configurations for page: Editor.py
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
[pages.Editor]
|
||||
page_title = "Editor"
|
||||
page_icon = "https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/favicon.png"
|
||||
sidebar_image_path = "https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/main_logo.png"
|
||||
@@ -0,0 +1,23 @@
|
||||
div.row-widget > button:first-child {
|
||||
border-radius: 25px;
|
||||
}
|
||||
|
||||
[class="st-emotion-cache-1e9n592"] {
|
||||
border-radius: 25px !important;
|
||||
}
|
||||
|
||||
.box-default {
|
||||
border: 0.5px solid #6a90e2;
|
||||
padding: 10px;
|
||||
margin: 10px;
|
||||
height: 280px;
|
||||
border-radius: 25px;
|
||||
}
|
||||
|
||||
.box-clicked {
|
||||
border: 2px solid #3367d6;
|
||||
padding: 10px;
|
||||
margin: 10px;
|
||||
height: 280px;
|
||||
border-radius: 25px;
|
||||
}
|
||||
@@ -0,0 +1,6 @@
|
||||
.element-container:has(#button-after) + div button {
|
||||
min-height: 100px;
|
||||
max-width: 500px;
|
||||
min-width: 500px;
|
||||
border-radius: 25px;
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
[data-testid="stSidebarNav"] {
|
||||
background-image: url("https://storage.googleapis.com/github-repo/generative-ai/sample-apps/accelerating-product-innovation/app_images/main_logo.png");
|
||||
background-repeat: no-repeat;
|
||||
background-position: 50% 10%;
|
||||
margin-top: 5%;
|
||||
background-size: 80%;
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
<html>
|
||||
<head>
|
||||
<title>Start Auto Download file</title>
|
||||
<script src="https://code.jquery.com/jquery-3.2.1.min.js"></script>
|
||||
<script>
|
||||
$(
|
||||
'<a href="data:application/zip;base64,{b64}" download="{download_filename}">',
|
||||
)[0].click();
|
||||
</script>
|
||||
</head>
|
||||
</html>
|
||||
@@ -0,0 +1,28 @@
|
||||
"""
|
||||
Entry page of streamlit application.
|
||||
"""
|
||||
|
||||
from app.pages_utils import setup
|
||||
from app.pages_utils.pages_config import PAGES_CFG
|
||||
from st_pages import show_pages_from_config
|
||||
import streamlit as st
|
||||
|
||||
# Initialize session state if not already initialized
|
||||
if "initialize_session_state" not in st.session_state:
|
||||
st.session_state.initialize_session_state = False
|
||||
|
||||
# Initialize session state if not already initialized
|
||||
if st.session_state.initialize_session_state is False:
|
||||
setup.initialize_all_session_state()
|
||||
st.session_state.initialize_session_state = True
|
||||
|
||||
# get the page configuration for the home page
|
||||
page_cfg = PAGES_CFG["home"]
|
||||
setup.page_setup(page_cfg)
|
||||
|
||||
|
||||
show_pages_from_config()
|
||||
|
||||
st.image(page_cfg["home_img_1"])
|
||||
st.divider()
|
||||
st.image(page_cfg["home_img_2"])
|
||||
@@ -0,0 +1,92 @@
|
||||
"""This module manages the image editing page within a Streamlit application.
|
||||
It provides the following features:
|
||||
|
||||
* Image Upload Handling: Processes and stores uploaded images.
|
||||
* Interactive Image Editor:
|
||||
* Offers an interface for modifying images (drawing, background editing,
|
||||
etc.).
|
||||
* Processes and combines foreground edits with the background image.
|
||||
* Draft Saving: Enables saving edited images as drafts, replacing the original
|
||||
draft versions.
|
||||
* Image Suggestion Generation:
|
||||
* Accepts text prompts to generate variations of the edited image.
|
||||
* Displays the generated image suggestions.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
import streamlit as st
|
||||
from PIL import Image
|
||||
from app.pages_utils import setup
|
||||
from app.pages_utils.edit_image import (
|
||||
generate_suggested_images,
|
||||
handle_image_upload,
|
||||
initialize_edit_page_state,
|
||||
process_foreground_image,
|
||||
render_suggested_images,
|
||||
save_draft_image,
|
||||
)
|
||||
from app.pages_utils.editor_ui import ImageEditor
|
||||
from app.pages_utils.pages_config import PAGES_CFG
|
||||
|
||||
# Get the configuration for the edit page
|
||||
page_cfg = PAGES_CFG["Editor"]
|
||||
|
||||
# Set up the general state of the app if uninitialized.
|
||||
setup.page_setup(page_cfg)
|
||||
|
||||
# Initialize the state of the edit page
|
||||
initialize_edit_page_state()
|
||||
|
||||
# Set up logging
|
||||
logging.basicConfig(format="%(levelname)s:%(message)s", level=logging.DEBUG)
|
||||
|
||||
# Check if the user has uploaded an image
|
||||
if st.session_state.uploaded_img is True:
|
||||
handle_image_upload()
|
||||
|
||||
# Check if the user has started editing the image
|
||||
if st.session_state.start_editing is None or st.session_state.start_editing is True:
|
||||
# Initialize Editor
|
||||
image_editor = ImageEditor()
|
||||
|
||||
# Display the image editor UI
|
||||
canvas_result, bg_image, image_bytes = image_editor.display_ui()
|
||||
background = Image.new("RGB", bg_image.size)
|
||||
|
||||
# Display save button only if draft elements exist
|
||||
# Function of save button: replace the original image in drafts with
|
||||
# edited image
|
||||
if st.session_state.draft_elements is not None:
|
||||
if st.button("Save"):
|
||||
row = st.session_state.image_edit_row
|
||||
col = st.session_state.image_edit_col
|
||||
image = bg_image
|
||||
drafts = st.session_state.draft_elements
|
||||
|
||||
save_draft_image(row, col, image, drafts)
|
||||
|
||||
# Check if drawing exists on the canvas (i.e., not blank)
|
||||
if canvas_result.image_data is not None and canvas_result.image_data.any():
|
||||
# Convert canvas data to a PIL Image object
|
||||
foreground = Image.fromarray(canvas_result.image_data)
|
||||
# Call image processing function, using canvas drawing as foreground
|
||||
processed_image_bytes = process_foreground_image(
|
||||
foreground_image=foreground,
|
||||
background_image=background,
|
||||
bg_editing=st.session_state.bg_editing,
|
||||
)
|
||||
# Store the processed image data for further use
|
||||
st.session_state.mask_image = processed_image_bytes
|
||||
|
||||
# If the prompt to generate edited images has been submitted
|
||||
if st.session_state.generate_images is True:
|
||||
generate_suggested_images(
|
||||
st.session_state.image_prompt,
|
||||
image_bytes,
|
||||
st.session_state.mask_image,
|
||||
)
|
||||
|
||||
# If the image generation is complete, render suggestion on ui
|
||||
if st.session_state.suggested_images is not None:
|
||||
render_suggested_images(st.session_state.suggested_images)
|
||||
@@ -0,0 +1,165 @@
|
||||
"""
|
||||
This module manages the 'Generations' page in the Streamlit application,
|
||||
with a focus on guiding users through product generation. This module:
|
||||
|
||||
* Project Management: Displays the project selected by the user.
|
||||
* Prompt-Based Generation:
|
||||
* Provides a form for users to specify product characteristics.
|
||||
* Generates product features based on the provided input.
|
||||
* Content Drafts:
|
||||
* Creates and displays new product ideas based on generated features.
|
||||
* Allows users to modify their selected features and content drafts.
|
||||
* Export Options: Enables downloading generated content and creating email
|
||||
copies.
|
||||
* Image Editing: Facilitates redirection to an image editing page.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
|
||||
from app.pages_utils import setup
|
||||
from app.pages_utils.downloads import download_content, download_file
|
||||
from app.pages_utils.draft_generation import ProductDrafts
|
||||
from app.pages_utils.pages_config import PAGES_CFG
|
||||
from app.pages_utils.product_features import (
|
||||
generate_formatted_response,
|
||||
modify_selection,
|
||||
render_features,
|
||||
)
|
||||
from app.pages_utils.product_gen import (
|
||||
build_prompt_form,
|
||||
handle_content_generation,
|
||||
update_generation_state,
|
||||
)
|
||||
import streamlit as st
|
||||
|
||||
|
||||
@st.cache_data
|
||||
def get_prod_gen_img() -> None:
|
||||
"""
|
||||
This function loads an image from a file, displays, and caches it.
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
# Display the top image for this page.
|
||||
page_images = [page_cfg["prod_gen_img"]]
|
||||
for page_image in page_images:
|
||||
st.image(page_image)
|
||||
|
||||
|
||||
def initialize_prod_gen() -> None:
|
||||
"""
|
||||
This function initializes the session state for the product generation
|
||||
page.
|
||||
|
||||
Args:
|
||||
None
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
# All images generated on this page have prefix 'gen_image'
|
||||
st.session_state.image_file_prefix = "gen_image"
|
||||
st.session_state.image_to_edit = -1 # No image is being edited.
|
||||
st.session_state.text_to_edit = -1 # Text is not being edited
|
||||
# Tracks whether image suggestions have been generated (on edit image).
|
||||
st.session_state.suggested_images = None
|
||||
st.session_state.generate_images = (
|
||||
False # Tracks whether images for product ideas have been generated.
|
||||
)
|
||||
# Display header images
|
||||
get_prod_gen_img()
|
||||
|
||||
|
||||
# Initialize page config.
|
||||
page_cfg = PAGES_CFG["3_Generations"]
|
||||
|
||||
setup.page_setup(page_cfg)
|
||||
|
||||
# Set product generation states.
|
||||
initialize_prod_gen()
|
||||
|
||||
# Page styles
|
||||
setup.load_css("app/css/prod_gen_styles.css")
|
||||
# logging initialization
|
||||
logging.basicConfig(format="%(levelname)s:%(message)s", level=logging.DEBUG)
|
||||
|
||||
# page title
|
||||
st.write(
|
||||
"""This page provides a step-by-step guide to generating products with
|
||||
desired characteristics."""
|
||||
)
|
||||
|
||||
# Display the project selected by the user
|
||||
setup.display_projects()
|
||||
|
||||
# Product generate form
|
||||
generate_btn = build_prompt_form()
|
||||
if generate_btn:
|
||||
update_generation_state()
|
||||
|
||||
# After form submission, generate product features
|
||||
if st.session_state.features_generated is True:
|
||||
if st.session_state.generated_response is None:
|
||||
st.session_state.generated_response = generate_formatted_response(
|
||||
st.session_state.selected_prompt
|
||||
)
|
||||
st.session_state.generated_points = None
|
||||
|
||||
# feature container
|
||||
features = st.empty()
|
||||
# Generated features to be displayed only if product content is not generated
|
||||
if (
|
||||
st.session_state.generated_response is not None
|
||||
and st.session_state.features_generated is True
|
||||
):
|
||||
if st.session_state.content_generated is False:
|
||||
# Display features on ui
|
||||
render_features(features)
|
||||
|
||||
# Columns for four buttons for product content
|
||||
content_gen_buttons = st.columns([10, 4, 10, 4, 10, 4, 10])
|
||||
|
||||
# product content container
|
||||
content = st.empty()
|
||||
|
||||
# Generate Button
|
||||
with content_gen_buttons[0]:
|
||||
# Get content corresponding to the features
|
||||
if st.button("Generate Content", type="primary"):
|
||||
asyncio.run(handle_content_generation(features))
|
||||
|
||||
# Display the generated content drafts
|
||||
product_drafts = ProductDrafts()
|
||||
product_drafts.display_drafts()
|
||||
|
||||
if st.session_state.create_product is True:
|
||||
# Modify Button
|
||||
with content_gen_buttons[2]:
|
||||
modify_btn = st.button("Modify Selection", type="primary")
|
||||
|
||||
# Redisplay the product drafts if content is being modified
|
||||
if modify_btn:
|
||||
modify_selection(content)
|
||||
|
||||
# Email download Button
|
||||
with content_gen_buttons[4]:
|
||||
email_dl_btn = st.button(
|
||||
"Generate Email Copy",
|
||||
on_click=download_file,
|
||||
type="primary",
|
||||
)
|
||||
|
||||
# Download Content Button
|
||||
with content_gen_buttons[6]:
|
||||
if st.session_state.product_content is not None:
|
||||
export_btn = st.button(
|
||||
"Export Content",
|
||||
on_click=download_content,
|
||||
type="primary",
|
||||
)
|
||||
|
||||
# If user clicks edit image, redirect to edit page
|
||||
if st.session_state.image_to_edit != -1 or st.session_state.generate_images is True:
|
||||
st.switch_page("pages/edit_image.py")
|
||||
@@ -0,0 +1,228 @@
|
||||
"""This module manages the 'Insights' page of the Streamlit application.
|
||||
It provides the following functionality:
|
||||
|
||||
* Project Management: Displays currently selected project.
|
||||
* Data Handling:
|
||||
* Checks for uploaded files in the specified project category.
|
||||
* Loads File Data (if uploaded files exist).
|
||||
* Question & Answer (QA):
|
||||
* Offers suggested questions for insights (requires uploaded data)
|
||||
through Retrieval Augmented Generation.
|
||||
* Enables users to input their own questions.
|
||||
* Generates answers to questions based on uploaded data, including
|
||||
references.
|
||||
* Follow-up Questions: Suggests follow-up questions based on previous queries.
|
||||
"""
|
||||
|
||||
import streamlit as st
|
||||
from app.pages_utils import insights, setup
|
||||
from app.pages_utils.pages_config import PAGES_CFG
|
||||
|
||||
# Get page configuration from config file
|
||||
page_cfg = PAGES_CFG["2_Marketing_Insights"]
|
||||
setup.page_setup(page_cfg)
|
||||
|
||||
# Initialize temporary suggestions if not already initialized
|
||||
if "temp_suggestions" not in st.session_state:
|
||||
st.session_state.temp_suggestions = None
|
||||
|
||||
|
||||
# Cache the function to get insights images
|
||||
@st.cache_data
|
||||
def get_insights_img() -> None:
|
||||
"""Loads, displays and cache header image for insights page.
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
page_images = [page_cfg["prod_insights_1"], page_cfg["prod_insights_2"]]
|
||||
for page_image in page_images:
|
||||
st.image(page_image)
|
||||
st.divider()
|
||||
|
||||
|
||||
# Display header images
|
||||
get_insights_img()
|
||||
|
||||
# Display projects
|
||||
setup.display_projects()
|
||||
|
||||
# Check if DataFrame is empty
|
||||
if st.session_state.embeddings_df is None or st.session_state.embeddings_df.empty:
|
||||
with st.spinner("Fetching Uploaded Files..."):
|
||||
embeddings_df = insights.get_stored_embeddings_as_df()
|
||||
st.session_state.embeddings_df = embeddings_df
|
||||
|
||||
|
||||
# Function to display suggestion box
|
||||
def display_suggestion_box(key: str, suggestion_num: int) -> None:
|
||||
"""Styles and displays the suggestion for insight generation.
|
||||
|
||||
Args:
|
||||
key: Unique key for suggestion.
|
||||
suggestion_num: Current suggestion number.
|
||||
"""
|
||||
# Apply custom styling to suggestion box
|
||||
setup.load_css("app/css/prod_insights_styles.css")
|
||||
# Display suggestion button
|
||||
if st.button(
|
||||
st.session_state.insights_suggestion[suggestion_num],
|
||||
key=key,
|
||||
):
|
||||
# Update session state with selected suggestion
|
||||
st.session_state.insights_placeholder = st.session_state.insights_suggestion[
|
||||
suggestion_num
|
||||
]
|
||||
# Set flag to generate RAG answers
|
||||
st.session_state.rag_answers_gen = True
|
||||
|
||||
|
||||
# Display suggested questions
|
||||
st.divider()
|
||||
st.write("**SUGGESTED QUESTIONS:**")
|
||||
|
||||
# Check if DataFrame is empty
|
||||
if st.session_state.embeddings_df is None or st.session_state.embeddings_df.empty:
|
||||
# Display error message
|
||||
st.error(
|
||||
"Add files in "
|
||||
+ st.session_state.product_category
|
||||
+ " file storage to get suggested questions",
|
||||
icon="🚨",
|
||||
)
|
||||
else:
|
||||
# Loop until suggestions are loaded or try limit is reached
|
||||
for try_var in range(3):
|
||||
if (
|
||||
st.session_state.suggestion_first_time
|
||||
and st.session_state.insights_suggestion is None
|
||||
):
|
||||
with st.spinner("Loading suggestion..."):
|
||||
insights.get_suggestions("insights_suggestion")
|
||||
# Check if number of suggestions is less than 4
|
||||
if (
|
||||
st.session_state.insights_suggestion is not None
|
||||
and len(st.session_state.insights_suggestion) < 4 # type: ignore
|
||||
):
|
||||
insights.get_suggestions("insights_suggestion") # type: ignore
|
||||
|
||||
# Check if suggestions are loaded
|
||||
if (
|
||||
st.session_state.insights_suggestion is None
|
||||
or len(st.session_state.insights_suggestion) < 4
|
||||
):
|
||||
st.write(st.session_state.insights_suggestion)
|
||||
# Display error message
|
||||
st.error("Sorry couldn't load suggestion")
|
||||
else:
|
||||
# Clear previous suggestions
|
||||
st.empty()
|
||||
|
||||
# Create columns for suggestion boxes
|
||||
suggestion_col1 = st.columns(2)
|
||||
|
||||
# Display suggestion boxes in first column
|
||||
with suggestion_col1[0]:
|
||||
display_suggestion_box("001", 0)
|
||||
with suggestion_col1[1]:
|
||||
display_suggestion_box("002", 1)
|
||||
|
||||
# Create columns for suggestion boxes
|
||||
suggestion_col2 = st.columns(2)
|
||||
|
||||
# Display suggestion boxes in second column
|
||||
with suggestion_col2[0]:
|
||||
display_suggestion_box("003", 2)
|
||||
with suggestion_col2[1]:
|
||||
display_suggestion_box("004", 3)
|
||||
|
||||
# Display divider
|
||||
st.divider()
|
||||
|
||||
# Get search term from text area
|
||||
search_term = st.text_area(
|
||||
"",
|
||||
key="2",
|
||||
value=st.session_state.insights_placeholder,
|
||||
placeholder="Select a suggestion or type your question here",
|
||||
)
|
||||
|
||||
# Display search button
|
||||
if st.button("Search", type="primary"):
|
||||
# Set flag to generate RAG answers
|
||||
st.session_state.rag_answers_gen = True
|
||||
|
||||
# Check if RAG answers should be generated
|
||||
if st.session_state.rag_answers_gen:
|
||||
# Check if DataFrame is empty
|
||||
if st.session_state.embeddings_df is None or st.session_state.embeddings_df.empty:
|
||||
# Display error message
|
||||
st.error(
|
||||
"Add files in " + st.session_state.product_category + " file storage",
|
||||
icon="🚨",
|
||||
)
|
||||
# Check if search term is empty
|
||||
elif search_term == "":
|
||||
# Display error message
|
||||
st.error(
|
||||
"Write the query to get the answer",
|
||||
icon="🚨",
|
||||
)
|
||||
else:
|
||||
# Clear previous results
|
||||
st.empty()
|
||||
|
||||
# Update session state with search term
|
||||
st.session_state.rag_search_term = search_term
|
||||
|
||||
# Generate RAG answers and references
|
||||
with st.spinner("Loading answer..."):
|
||||
st.session_state.rag_search_term = search_term
|
||||
(
|
||||
st.session_state.rag_answer,
|
||||
st.session_state.rag_answer_references,
|
||||
) = insights.generate_insights_search_result(
|
||||
st.session_state.rag_search_term
|
||||
)
|
||||
|
||||
# Get new suggestions
|
||||
with st.spinner("Getting new Suggestions"):
|
||||
insights.get_suggestions("temp_suggestions")
|
||||
|
||||
# Check if RAG answer and references are available
|
||||
if (
|
||||
st.session_state.rag_answer is not None
|
||||
and st.session_state.rag_answer_references is not None
|
||||
):
|
||||
# Display RAG answer
|
||||
st.write(st.session_state.rag_answer)
|
||||
st.write()
|
||||
|
||||
# Display RAG answer references
|
||||
st.write("**REFERENCES**")
|
||||
st.write(st.session_state.rag_answer_references)
|
||||
|
||||
# Reset RAG answers generation flag
|
||||
st.session_state.rag_answers_gen = False
|
||||
st.session_state.suggestion_first_time = 0
|
||||
|
||||
# Check if temporary suggestions are available
|
||||
if st.session_state.temp_suggestions is not None:
|
||||
# Display divider
|
||||
st.divider()
|
||||
|
||||
# Display follow-up questions
|
||||
st.write("**Follow up questions**")
|
||||
|
||||
# Display follow-up question buttons
|
||||
for suggestion in st.session_state.temp_suggestions:
|
||||
if st.button(
|
||||
suggestion,
|
||||
key=f"{suggestion} {st.session_state.rag_search_term}",
|
||||
):
|
||||
# Update session state with selected suggestion
|
||||
st.session_state.insights_placeholder = suggestion
|
||||
# Set flag to generate RAG answers
|
||||
st.session_state.rag_answers_gen = True
|
||||
# Rerun the app
|
||||
st.rerun()
|
||||
@@ -0,0 +1,3 @@
|
||||
streamlit
|
||||
pillow
|
||||
google-cloud-storage==2.19.0
|
||||
@@ -0,0 +1,192 @@
|
||||
"""
|
||||
This module manages the "Resources" page of the Streamlit application. Key
|
||||
functionalities include:
|
||||
|
||||
* Project Management:
|
||||
* Displays existing projects.
|
||||
* Allows users to add new project categories.
|
||||
* File Management:
|
||||
* Enables file uploads (txt, docx, pdf, csv).
|
||||
* Handles conversion and storage of uploaded files.
|
||||
* Lists stored project files.
|
||||
* Provides download and delete options for stored files.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
from app.pages_utils import project, resources_store_embeddings, setup
|
||||
from app.pages_utils.pages_config import GLOBAL_CFG, PAGES_CFG
|
||||
from google.cloud import storage
|
||||
import streamlit as st
|
||||
|
||||
# Get the page configuration from the config file
|
||||
page_cfg = PAGES_CFG["1_Resources"]
|
||||
setup.page_setup(page_cfg)
|
||||
|
||||
|
||||
PROJECT_ID = os.getenv("PROJECT_ID")
|
||||
LOCATION = os.getenv("LOCATION")
|
||||
|
||||
# Define storage bucket
|
||||
storage_client = storage.Client(project=PROJECT_ID)
|
||||
bucket = storage_client.bucket(GLOBAL_CFG["bucket_name"])
|
||||
|
||||
|
||||
# Initialize project form submission state if not already initialized
|
||||
if "project_form_submitted" not in st.session_state:
|
||||
st.session_state.project_form_submitted = False
|
||||
|
||||
|
||||
@st.cache_data
|
||||
def get_resources_img() -> None:
|
||||
"""
|
||||
Loads, displays and caches the resources header image.
|
||||
|
||||
Returns: None.
|
||||
"""
|
||||
# Get the file path of the resources image
|
||||
page_images = [page_cfg["resources_img"]]
|
||||
for page_image in page_images:
|
||||
st.image(page_image)
|
||||
|
||||
|
||||
# Display header image.
|
||||
get_resources_img()
|
||||
|
||||
# Create a container for the screen
|
||||
screen = st.container()
|
||||
|
||||
# Create a form for the resources page
|
||||
with st.form(key="resources_form", clear_on_submit=True):
|
||||
# Display the projects
|
||||
setup.display_projects()
|
||||
|
||||
# Add a text input field for adding a new project category
|
||||
st.session_state.new_product_category_added = st.text_input(
|
||||
"",
|
||||
key="000",
|
||||
placeholder="Add a new project",
|
||||
)
|
||||
|
||||
# Add a file uploader for uploading files
|
||||
st.session_state.uploaded_files = st.file_uploader(
|
||||
"",
|
||||
type=["txt", "docx", "pdf", "csv"],
|
||||
accept_multiple_files=True,
|
||||
key="09",
|
||||
)
|
||||
|
||||
# Add a submit button to the form
|
||||
submitted = st.form_submit_button("Submit", type="primary")
|
||||
|
||||
# Check if the form was submitted
|
||||
if submitted:
|
||||
st.session_state.project_form_submitted = True
|
||||
# Check if a new project category was added
|
||||
if (
|
||||
st.session_state.new_product_category_added is not None
|
||||
and st.session_state.new_product_category_added != ""
|
||||
):
|
||||
# Update the product category list
|
||||
st.session_state.product_category = st.session_state.new_product_category_added
|
||||
st.session_state.product_categories = [
|
||||
st.session_state.new_product_category_added
|
||||
] + st.session_state.product_categories
|
||||
|
||||
# Update the projects in GCS
|
||||
project_list_blob = bucket.blob("project_list.txt")
|
||||
project_list_blob.upload_from_string(
|
||||
json.dumps(st.session_state.product_categories)
|
||||
)
|
||||
|
||||
# Reset the new project category field
|
||||
st.session_state.new_product_category_added = None
|
||||
|
||||
# Check if files were uploaded
|
||||
if st.session_state.uploaded_files is not None:
|
||||
# Convert the uploaded files to data packets and upload them to GCS
|
||||
for uploaded_file in st.session_state.uploaded_files:
|
||||
resources_store_embeddings.create_and_store_embeddings(uploaded_file)
|
||||
|
||||
# Check if files were uploaded
|
||||
if st.session_state.uploaded_files is not None:
|
||||
# Convert the uploaded files to data packets and upload them to GCS
|
||||
for uploaded_file in st.session_state.uploaded_files:
|
||||
resources_store_embeddings.create_and_store_embeddings(uploaded_file)
|
||||
|
||||
|
||||
# Check if the project form was submitted and the file upload is complete
|
||||
if st.session_state.project_form_submitted is True:
|
||||
# Create columns for the project heading and delete button
|
||||
project_heading = st.columns([4, 1])
|
||||
|
||||
# Display the project category heading
|
||||
with project_heading[0]:
|
||||
st.markdown(
|
||||
f"""<h3 style = 'text-align: center; color: #6a90e2;'>
|
||||
{st.session_state.product_category}</h3>""",
|
||||
unsafe_allow_html=True,
|
||||
)
|
||||
|
||||
# Add a delete button for the project
|
||||
with project_heading[1]:
|
||||
if st.button("Delete this project", type="primary"):
|
||||
# Display a spinner while deleting the project
|
||||
with st.spinner("Deleting Project..."):
|
||||
# Delete the project from GCS
|
||||
project.delete_project_from_gcs()
|
||||
|
||||
# List the PDF files in the GCS bucket
|
||||
files = project.list_pdf_files_gcs()
|
||||
|
||||
# Get the length of the product category name
|
||||
len_prod_cat = len(st.session_state.product_category) + 1
|
||||
|
||||
# Display the files in a spinner
|
||||
with st.spinner("Fetching Files"):
|
||||
# Set the border style for the file list items
|
||||
BORDER_STYLE = "border: 2px solid black; padding: 10px;"
|
||||
|
||||
# Iterate over the files
|
||||
for color_counter, file in enumerate(files):
|
||||
# Create columns for the file name, download button, and
|
||||
# delete button
|
||||
list_files_columns = st.columns([15, 1, 1])
|
||||
|
||||
# Set a color counter to alternate the background color of the file.
|
||||
with list_files_columns[0]:
|
||||
if color_counter % 2 == 0:
|
||||
BACKGROUND_COLOR = "#e6f2ff"
|
||||
else:
|
||||
BACKGROUND_COLOR = "white"
|
||||
st.write(
|
||||
f"""<div style=
|
||||
'background-color: {BACKGROUND_COLOR};
|
||||
padding: 10px;
|
||||
margin-bottom: 0px;
|
||||
margin-top: 0px;
|
||||
border-radius:10px;'>{file[0][len_prod_cat:]}</div>""",
|
||||
unsafe_allow_html=True,
|
||||
)
|
||||
|
||||
# Add a download button for the file
|
||||
file_content_blob = bucket.blob(
|
||||
f"{st.session_state.product_category}/{file[0][len_prod_cat:]}"
|
||||
)
|
||||
file_contents = file_content_blob.download_as_string()
|
||||
with list_files_columns[1]:
|
||||
st.download_button(
|
||||
label=":arrow_down:",
|
||||
data=file_contents,
|
||||
file_name=file[0][len_prod_cat:],
|
||||
mime=file[1],
|
||||
)
|
||||
|
||||
# Add a delete button for the file
|
||||
with list_files_columns[2]:
|
||||
if st.button(
|
||||
":x:",
|
||||
key=file[0][len_prod_cat:],
|
||||
):
|
||||
project.delete_file_from_gcs(file_name=file[0][len_prod_cat:])
|
||||
@@ -0,0 +1,199 @@
|
||||
"""
|
||||
This module provides functions for downloading generated content (emails,
|
||||
product content)
|
||||
as zip archives.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import io
|
||||
import logging
|
||||
from typing import Any
|
||||
import zipfile
|
||||
|
||||
from app.pages_utils.export_content_pdf import create_content_pdf, create_email_pdf
|
||||
from app.pages_utils.get_llm_response import generate_gemini
|
||||
from app.pages_utils.imagen import image_generation
|
||||
from dotenv import load_dotenv
|
||||
import streamlit as st
|
||||
import streamlit.components.v1 as components
|
||||
|
||||
load_dotenv()
|
||||
|
||||
logging.basicConfig(format="%(levelname)s:%(message)s", level=logging.DEBUG)
|
||||
st.session_state.image_file_prefix = "email_image"
|
||||
|
||||
|
||||
def generate_email(prompt: str, title: str) -> None:
|
||||
"""Generates an email PDF with the given prompt and title.
|
||||
|
||||
Args:
|
||||
prompt (str): The prompt for the email text.
|
||||
title (str): The title of the email.
|
||||
"""
|
||||
# Create prompt for email generation for given product idea.
|
||||
email_prompt = f"""Write an email introducing the concept for a new
|
||||
product, {prompt}. Write the benefits for different demographics,
|
||||
skin types, genders, etc., too.
|
||||
The email should strictly share the concept with the innovation team.
|
||||
The email should strictly not announce launch, only the concept.
|
||||
Keep the email brief."""
|
||||
|
||||
# Generate Email content.
|
||||
email_text = generate_gemini(email_prompt)
|
||||
st.session_state.email_text = email_text # update state.
|
||||
|
||||
# Generate corresponding image for email copy.
|
||||
image_generation(
|
||||
f"""Generate a beautiful image of a {st.session_state.product_category}
|
||||
in an aesthetic background. Image should be suitable for advertising.
|
||||
Content should be written on packaging in English.""",
|
||||
1,
|
||||
"1:1",
|
||||
"email_image",
|
||||
)
|
||||
|
||||
# Generate pdf containing the email content and image.
|
||||
create_email_pdf(
|
||||
title,
|
||||
email_text.replace("**", ""),
|
||||
f"email_copy_0_{title}",
|
||||
"email_image.png",
|
||||
)
|
||||
st.session_state.email_files.append(f"email_copy_0_{title}.pdf")
|
||||
|
||||
|
||||
def download_button(object_to_download: bytes, download_filename: str) -> str:
|
||||
"""Generates a download link for the given object.
|
||||
|
||||
Args:
|
||||
object_to_download (bytes or str): The object to download.
|
||||
download_filename (str): The filename of the downloaded object.
|
||||
|
||||
Returns:
|
||||
str: The HTML code for the download link.
|
||||
"""
|
||||
# Create a BytesIO object to hold the zip file content
|
||||
zip_buffer = io.BytesIO()
|
||||
|
||||
with zipfile.ZipFile(zip_buffer, "a", zipfile.ZIP_DEFLATED) as zip_file:
|
||||
if isinstance(object_to_download, bytes):
|
||||
# If it's already bytes (e.g., binary data), add it to the zip file
|
||||
zip_file.writestr(download_filename, object_to_download)
|
||||
else:
|
||||
# If it's not bytes, handle accordingly (modify as needed)
|
||||
raise ValueError("Unsupported type for object_to_download")
|
||||
|
||||
# Get the BytesIO object's content as bytes
|
||||
zip_content = zip_buffer.getvalue()
|
||||
|
||||
# Encode the zip content in base64
|
||||
b64 = base64.b64encode(zip_content).decode()
|
||||
|
||||
# Read the HTML template file
|
||||
with open("app/download_link.html", encoding="utf8") as f:
|
||||
html_template = f.read()
|
||||
|
||||
# Replace placeholders in the HTML template
|
||||
html_link = html_template.replace("{b64}", b64)
|
||||
html_link = html_link.replace("{download_filename}", download_filename)
|
||||
|
||||
return html_link
|
||||
|
||||
|
||||
def create_zip_buffer(filenames: list[str]) -> io.BytesIO:
|
||||
"""Creates a BytesIO object containing a zip file of the specified files.
|
||||
|
||||
Args:
|
||||
filenames: A list of filenames to include in the zip archive.
|
||||
|
||||
Returns:
|
||||
An io.BytesIO object representing the zip file in memory.
|
||||
"""
|
||||
zip_buffer = io.BytesIO()
|
||||
|
||||
with zipfile.ZipFile(zip_buffer, "a", zipfile.ZIP_DEFLATED) as zip_file:
|
||||
for filename in filenames:
|
||||
with open(f"./{filename}", "rb") as pdf_file:
|
||||
zip_file.writestr(filename, pdf_file.read())
|
||||
return zip_buffer
|
||||
|
||||
|
||||
def load_product_lists() -> tuple[list[list[dict[str, Any]]], list[str]]:
|
||||
"""
|
||||
Creates copies of Product titles and content to be imported.
|
||||
Returns:
|
||||
Tuple containing two lists - Product Content and titles.
|
||||
"""
|
||||
# Create copies to avoid modifying session data
|
||||
prod_content = st.session_state.draft_elements.copy()
|
||||
titles = st.session_state.selected_titles.copy()
|
||||
|
||||
# Handle the case of multiple titles including assorted content
|
||||
if len(st.session_state.selected_titles) > 1:
|
||||
prod_content.append(st.session_state.assorted_prod_content)
|
||||
titles.append(st.session_state.assorted_prod_title)
|
||||
else:
|
||||
prod_content.append("")
|
||||
|
||||
return prod_content, titles
|
||||
|
||||
|
||||
def download_file() -> None:
|
||||
"""Downloads the generated email files as a zip archive."""
|
||||
|
||||
with st.spinner("Downloading Email files ..."):
|
||||
st.session_state.email_gen = True
|
||||
|
||||
if st.session_state.draft_elements is not None:
|
||||
prod_content, titles = load_product_lists()
|
||||
# Prepare file list for the zip file
|
||||
filenames = []
|
||||
|
||||
# Variable to store the name of email_file
|
||||
email_file_title = st.session_state.assorted_prod_title
|
||||
|
||||
# Logic to generate content for email files.
|
||||
for i, title in enumerate(titles):
|
||||
st.session_state.email_files = []
|
||||
if st.session_state.email_gen:
|
||||
generate_email(prod_content[i][0]["text"], title)
|
||||
|
||||
# Generate a single file for each title
|
||||
filename = f"{st.session_state.email_files[0]}"
|
||||
filenames.append(filename)
|
||||
|
||||
# Create the zip file in memory.
|
||||
zip_buffer = create_zip_buffer(filenames)
|
||||
|
||||
# Provide download button with appropriate filename
|
||||
components.html(
|
||||
download_button(zip_buffer.getvalue(), f"email_{email_file_title}.zip"),
|
||||
height=0,
|
||||
)
|
||||
st.success("Email Copies Downloaded")
|
||||
|
||||
|
||||
def download_content() -> None:
|
||||
"""Downloads the generated content as a zip archive."""
|
||||
|
||||
with st.spinner("Creating Content pdf"):
|
||||
prod_content, titles = load_product_lists()
|
||||
|
||||
# Call the function to generate content PDFs
|
||||
create_content_pdf(prod_content, titles)
|
||||
|
||||
# Create the zip archive
|
||||
filenames = []
|
||||
|
||||
# Generate filenames
|
||||
for i in range(len(titles)):
|
||||
filenames.append(f"content_{i}.pdf")
|
||||
|
||||
zip_buffer = create_zip_buffer(filenames)
|
||||
|
||||
# Prepare download button with a dynamic filename
|
||||
components.html(
|
||||
download_button(zip_buffer.getvalue(), f"content_{titles[i]}.zip"),
|
||||
height=0,
|
||||
)
|
||||
st.success("Downloaded Content Zip.")
|
||||
+141
@@ -0,0 +1,141 @@
|
||||
"""
|
||||
This module defines the 'ProductDrafts' class, responsible
|
||||
for managing and displaying product content drafts.
|
||||
"""
|
||||
|
||||
from app.pages_utils.get_llm_response import generate_gemini
|
||||
import streamlit as st
|
||||
|
||||
|
||||
class ProductDrafts:
|
||||
"""
|
||||
Key functionalities include:
|
||||
|
||||
* Draft Organization: Arranges drafts based on their associated product
|
||||
titles.
|
||||
* Display: Renders drafts in an expandable format, including both image
|
||||
and text components.
|
||||
* Edit Integration: Provides buttons to trigger image editing and text
|
||||
regeneration for each draft.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
# Initialize session state if needed
|
||||
if "create_product" not in st.session_state:
|
||||
st.session_state.create_product = False
|
||||
|
||||
def display_drafts(self) -> None:
|
||||
"""
|
||||
Displays all product drafts if the 'create_product' flag is True in
|
||||
session state.
|
||||
"""
|
||||
|
||||
if st.session_state.create_product:
|
||||
for i, chosen_title in enumerate(st.session_state.chosen_titles):
|
||||
# Create an expandable section for each product title
|
||||
with st.expander(chosen_title):
|
||||
# Display a centered title heading
|
||||
st.markdown(
|
||||
f"""<h5 style = 'text-align: center; color: #6a90e2;'>
|
||||
{chosen_title}
|
||||
</h5>""",
|
||||
unsafe_allow_html=True,
|
||||
)
|
||||
|
||||
# Call the helper function to display drafts for this title
|
||||
self.display_draft_row(i)
|
||||
|
||||
def display_draft_row(self, title_index: int) -> None:
|
||||
"""
|
||||
Displays a single row of drafts for a given product title.
|
||||
|
||||
Args:
|
||||
title_index (int): The index of the product title in the
|
||||
st.session_state.chosen_titles list.
|
||||
"""
|
||||
|
||||
for j in range(st.session_state.num_drafts):
|
||||
img_col, text_col = st.columns(2) # Create two equal-width columns
|
||||
|
||||
with img_col:
|
||||
# Display the image for the current draft
|
||||
st.image(st.session_state.draft_elements[title_index][j]["img"])
|
||||
|
||||
# Call the function to handle image editing interactions
|
||||
self._handle_image_edit(title_index, j)
|
||||
|
||||
with text_col:
|
||||
# Display the text for the current draft
|
||||
st.write(st.session_state.draft_elements[title_index][j]["text"])
|
||||
st.session_state.text_edit_prompt = st.text_input(
|
||||
key="edit_text_prompt" + str(title_index) + str(j),
|
||||
placeholder="Write a query to edit the text",
|
||||
label="Write prompt to edit text",
|
||||
)
|
||||
if st.button(
|
||||
"Regenerate",
|
||||
key=f"""
|
||||
edit text{st.session_state.chosen_titles[title_index]}
|
||||
{st.session_state.num_drafts*title_index+j+1}
|
||||
""",
|
||||
type="primary",
|
||||
):
|
||||
# On button click begin content regeneration.
|
||||
st.session_state.regenerate_btn = True
|
||||
st.session_state.row = title_index # Title number being edited.
|
||||
# In case of multiple drafts, track the draft number
|
||||
# being edited.
|
||||
st.session_state.col = j
|
||||
st.session_state.text_to_edit = st.session_state.draft_elements[
|
||||
title_index
|
||||
][j][
|
||||
"text"
|
||||
] # Text content being edited.
|
||||
|
||||
# Update content
|
||||
with st.spinner("Updating Content..."):
|
||||
new_text_prompt = f"""Prompt: Based on the given query
|
||||
change the given context and give
|
||||
only the revised context.
|
||||
Query:
|
||||
{st.session_state.text_edit_prompt}
|
||||
Context:
|
||||
{st.session_state.text_to_edit} """
|
||||
# Generate new text.
|
||||
text = generate_gemini(new_text_prompt)
|
||||
# Update text post text regeneration.
|
||||
st.session_state.draft_elements[st.session_state.row][
|
||||
st.session_state.col
|
||||
][
|
||||
"text"
|
||||
] = text # Update draft contents.
|
||||
# Reset Text edit status to default.
|
||||
st.session_state.text_edit_prompt = None
|
||||
st.session_state.regenerate_btn = False
|
||||
st.session_state.update_text_btn = None
|
||||
# Reload page to display updated content
|
||||
st.rerun()
|
||||
|
||||
def _handle_image_edit(self, title_index: int, draft_index: int) -> None:
|
||||
"""
|
||||
Handles image edit button interactions for a specific draft.
|
||||
|
||||
Args:
|
||||
title_index (int): The index of the product title.
|
||||
draft_index (int): The index of the draft within the title.
|
||||
"""
|
||||
|
||||
# Construct a unique button key using title and draft index
|
||||
button_key = f"Edit Image - {title_index} - {draft_index}"
|
||||
|
||||
if st.button("Edit Image", key=button_key, type="primary"):
|
||||
# Store information in session state to track which image is
|
||||
# being edited
|
||||
st.session_state.image_edit_row = title_index
|
||||
st.session_state.image_edit_col = draft_index
|
||||
|
||||
# Calculate a unique index for the image and update the
|
||||
# session state.
|
||||
st.session_state.image_to_edit = (
|
||||
st.session_state.num_drafts * title_index + draft_index
|
||||
)
|
||||
@@ -0,0 +1,264 @@
|
||||
"""This module provides functions for image processing and session state
|
||||
management
|
||||
related to image editing. This module:
|
||||
* process_foreground_image():
|
||||
* Prepares a foreground image for merging with a background.
|
||||
* Optionally removes white regions for background editing.
|
||||
* initialize_edit_page_state(): Initializes session state for the
|
||||
image editing
|
||||
page and handles uploaded images.
|
||||
* handle_image_upload(): Manages the image upload process and updates session
|
||||
state.
|
||||
* save_draft_image(): Saves edited draft images and facilitates a return to
|
||||
the product
|
||||
generation page.
|
||||
* generate_suggested_images():
|
||||
* Generates image variations based on text prompts, an initial image, and
|
||||
an optional mask.
|
||||
* render_suggested_images():
|
||||
* Displays generated suggestions in a grid.
|
||||
* Provides "Edit" and "Download" buttons for each suggestion.
|
||||
* _handle_edit_suggestion():
|
||||
* Handles the logic for editing a selected suggestion.
|
||||
"""
|
||||
|
||||
import io
|
||||
import logging
|
||||
|
||||
import PIL
|
||||
import streamlit as st
|
||||
from PIL import Image
|
||||
from app.pages_utils.imagen import predict_edit_image
|
||||
from vertexai.preview.vision_models import Image as vertex_image
|
||||
|
||||
logging.basicConfig(format="%(levelname)s:%(message)s", level=logging.DEBUG)
|
||||
|
||||
|
||||
def process_foreground_image(
|
||||
foreground_image: Image.Image,
|
||||
background_image: Image.Image,
|
||||
bg_editing: bool = False,
|
||||
) -> bytes:
|
||||
"""Processes a foreground image, optionally removing white regions,
|
||||
and prepares it for merging with a background image.
|
||||
|
||||
Args:
|
||||
foreground_image (Image.Image): The PIL Image object representing the
|
||||
foreground.
|
||||
background_image (Image.Image): The PIL Image object representing the
|
||||
background.
|
||||
bg_editing (bool, optional): If True, removes white regions from the
|
||||
foreground. Defaults to False.
|
||||
|
||||
Returns:
|
||||
bytes: The processed and merged image data as bytes.
|
||||
"""
|
||||
# Logic to edit background (invert mask)
|
||||
if bg_editing:
|
||||
# Get image foreground
|
||||
foreground_data = foreground_image.getdata()
|
||||
new_bytes = []
|
||||
# Invert pixels.
|
||||
for item in foreground_data:
|
||||
if item[0] == 255 and item[1] == 255 and item[2] == 255:
|
||||
new_bytes.append((255, 255, 255, 0))
|
||||
else:
|
||||
new_bytes.append((255, 255, 255, 1))
|
||||
|
||||
foreground_image.putdata(new_bytes)
|
||||
|
||||
# Resize and merge foreground with background
|
||||
resized_foreground = foreground_image.resize(background_image.size)
|
||||
merged_image = background_image.copy()
|
||||
merged_image.paste(resized_foreground, (0, 0), resized_foreground)
|
||||
|
||||
# Convert to bytes for storage
|
||||
with io.BytesIO() as buffer:
|
||||
merged_image.save(buffer, format="PNG")
|
||||
processed_image_bytes = buffer.getvalue()
|
||||
|
||||
return processed_image_bytes
|
||||
|
||||
|
||||
def initialize_edit_page_state() -> None:
|
||||
"""Initializes the session state for the image editing page.
|
||||
|
||||
This function checks if the session state has been initialized, and if not,
|
||||
it initializes it.
|
||||
It also checks if an image has been uploaded, and if so, it sets
|
||||
the session state accordingly.
|
||||
"""
|
||||
# Check which image file prefix points to the image to be edited
|
||||
if "image_to_edit" not in st.session_state or st.session_state.image_to_edit == -1:
|
||||
st.session_state.image_to_edit = (
|
||||
-1
|
||||
) # No image from generations is being edited.
|
||||
st.session_state.image_file_prefix = (
|
||||
"uploaded_image" # image prefix for editing uploaded image.
|
||||
)
|
||||
st.session_state.uploaded_img = True # Set image uploaded to true.
|
||||
else:
|
||||
st.session_state.uploaded_img = False # Generated image being edited.
|
||||
st.session_state.start_editing = True # Display canvas for editing.
|
||||
|
||||
|
||||
def handle_image_upload() -> None:
|
||||
"""Handles an image upload, saving the image and updating session state."""
|
||||
# Upload button
|
||||
uploaded_file = st.file_uploader("Upload an image")
|
||||
|
||||
# If image has been uploaded open and Save uploaded image.
|
||||
if uploaded_file is not None:
|
||||
try:
|
||||
image = Image.open(uploaded_file)
|
||||
filename = "uploaded_image0.png"
|
||||
image.save(filename)
|
||||
st.session_state.start_editing = True
|
||||
except (OSError, PIL.UnidentifiedImageError) as e:
|
||||
st.error(f"Error opening image: {e}")
|
||||
|
||||
|
||||
def save_draft_image(
|
||||
row: int, col: int, image: Image.Image, draft_elements: dict
|
||||
) -> None:
|
||||
"""Saves the draft image and updates session state for content editing.
|
||||
|
||||
Args:
|
||||
row (int): Row index of the image being edited.
|
||||
col (int): Column index of the image being edited.
|
||||
image (Image): The image object to be saved.
|
||||
draft_elements (dict): Dictionary holding the draft image elements.
|
||||
"""
|
||||
st.session_state.content_edited = True # Track whether image has been edited.
|
||||
draft_elements[row][col]["img"] = (
|
||||
image # Update the drafts to display updated image.
|
||||
)
|
||||
|
||||
# Calculate unique image filename and save image.
|
||||
image_num = st.session_state.num_drafts * row + col + 1
|
||||
image.save(f"gen_image{image_num}.png")
|
||||
|
||||
# Display the edited image on product generation image
|
||||
st.switch_page("pages/product_generation.py")
|
||||
|
||||
|
||||
def render_suggested_images(suggested_images: list[str]) -> None:
|
||||
"""Renders suggested images in a grid layout with "Edit" and "Download"
|
||||
buttons.
|
||||
|
||||
Args:
|
||||
suggested_images: A list of image paths or data to display as
|
||||
suggestions.
|
||||
"""
|
||||
# Set number of images to be displayed per row.
|
||||
num_suggestions_per_row = 3
|
||||
|
||||
# Iterate over image rows.
|
||||
for row_start in range(0, len(suggested_images), num_suggestions_per_row):
|
||||
# Create columns to display each image.
|
||||
suggestion_cols = st.columns(num_suggestions_per_row)
|
||||
for col_index, col in enumerate(suggestion_cols):
|
||||
# Calculate image index
|
||||
image_index = row_start + col_index
|
||||
with col:
|
||||
# Display image.
|
||||
st.image(suggested_images[image_index])
|
||||
# Add Edit image button for the current suggestion.
|
||||
if st.button(
|
||||
"Edit",
|
||||
key=f"edit suggestion {image_index}",
|
||||
type="primary",
|
||||
):
|
||||
_handle_edit_suggestion(image_index)
|
||||
# Add download button for current suggestion.
|
||||
image_data = suggested_images[image_index]
|
||||
st.download_button(
|
||||
label="Download",
|
||||
data=image_data,
|
||||
file_name=f"suggestion_{image_index}.png",
|
||||
mime="image/png",
|
||||
type="primary",
|
||||
)
|
||||
|
||||
|
||||
def _handle_edit_suggestion(image_index: int) -> None:
|
||||
"""Handles the logic for when the 'Edit' button is clicked.
|
||||
|
||||
Args:
|
||||
image_index (int): corresponding draft number of image being edited.
|
||||
"""
|
||||
# Get Byte data of the image.
|
||||
image_data = io.BytesIO(st.session_state.suggested_images[image_index])
|
||||
# Save image.
|
||||
with open("suggestion1.png", "wb") as f:
|
||||
f.write(image_data.getvalue())
|
||||
|
||||
# Update state.
|
||||
st.session_state.edit_suggestion = (
|
||||
True # Track whether a suggestion is being edited.
|
||||
)
|
||||
st.session_state.image_file_prefix = (
|
||||
"suggestion" # Image saved with prefix suggestion is being edited.
|
||||
)
|
||||
st.session_state.image_to_edit = 0 # Track which image is being edited
|
||||
st.session_state.mask_image = None # Reset Mask
|
||||
st.rerun()
|
||||
|
||||
|
||||
def save_image_for_editing(image_bytes: bytes, filename: str) -> None:
|
||||
"""Saves image for image editing by Imagen
|
||||
|
||||
Args:
|
||||
image_bytes (bytes): Image bytes for the image to be saved.
|
||||
filename (str): Name of the saved file.
|
||||
"""
|
||||
# Create a BytesIO object from the image bytes
|
||||
image_stream = io.BytesIO(image_bytes)
|
||||
|
||||
# Open the image using Pillow
|
||||
image = Image.open(image_stream)
|
||||
|
||||
# Save the image as a PNG
|
||||
image.save(f"{filename}.png", "PNG")
|
||||
|
||||
|
||||
def generate_suggested_images(
|
||||
image_prompt: str,
|
||||
image_bytes: io.BytesIO,
|
||||
mask_image: bytes,
|
||||
sample_count: int = 6,
|
||||
) -> None:
|
||||
"""Generates suggested images based on the provided prompt, image, and mask.
|
||||
Updates Streamlit session state with the generated images.
|
||||
|
||||
Args:
|
||||
image_prompt (str): Text prompt for image generation.
|
||||
image_bytes (BytesIO): Initial image data for editing.
|
||||
mask_image (bytes): Mask defining the region to
|
||||
edit (optional).
|
||||
"""
|
||||
save_image_for_editing(image_bytes.getvalue(), "image_to_edit")
|
||||
save_image_for_editing(mask_image, "mask")
|
||||
|
||||
st.session_state.suggested_images = [] # Clear previous suggestions
|
||||
|
||||
# Generated edit image results
|
||||
with st.spinner("Generating suggested images"):
|
||||
input_dict = {
|
||||
"prompt": image_prompt,
|
||||
"image": vertex_image.load_from_file("image_to_edit.png"),
|
||||
}
|
||||
|
||||
if mask_image:
|
||||
input_dict["mask"] = vertex_image.load_from_file("mask.png")
|
||||
|
||||
st.session_state["generated_image"] = predict_edit_image(
|
||||
instance_dict=input_dict,
|
||||
parameters={"sampleCount": sample_count},
|
||||
)
|
||||
|
||||
# Append newly generated suggestions to suggested images state key.
|
||||
for image_data in st.session_state.generated_image:
|
||||
st.session_state.suggested_images.append(image_data.__dict__["_loaded_bytes"])
|
||||
# End image generation.
|
||||
st.session_state.generate_images = False # Update generation state
|
||||
@@ -0,0 +1,119 @@
|
||||
"""
|
||||
This module defines the 'ImageEditor' class, providing an interactive image
|
||||
editing interface.
|
||||
"""
|
||||
|
||||
import io
|
||||
|
||||
from PIL import Image
|
||||
import streamlit as st
|
||||
from streamlit_drawable_canvas import st_canvas
|
||||
|
||||
|
||||
class ImageEditor:
|
||||
"""
|
||||
Functions include:
|
||||
* Drawing Tools: Offers tools for drawing masks (rectangles,
|
||||
free drawing, circles).
|
||||
* Customization: Allows control over stroke width and drawing mode.
|
||||
* Background Editing: Enables background modification (with masking).
|
||||
* Prompts: Facilitates image generation based on user-provided text
|
||||
prompts.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.stroke_width = 20 # Default stroke width
|
||||
self.stroke_color = "black"
|
||||
self.drawing_mode = "rect" # Default drawing mode
|
||||
self.realtime_update = True
|
||||
|
||||
def load_image(self, image_file: str) -> io.BytesIO:
|
||||
"""Load an image from a local file as BytesIO.
|
||||
Args:
|
||||
image_file: path to image file to be loaded.
|
||||
|
||||
Returns:
|
||||
Image bytes object.
|
||||
"""
|
||||
with open(image_file, "rb") as f:
|
||||
image_data = f.read()
|
||||
return io.BytesIO(image_data)
|
||||
|
||||
def display_ui(self) -> tuple[st_canvas, Image.Image, io.BytesIO]:
|
||||
"""Renders the main UI components of the image editor."""
|
||||
# - Load the image for editing
|
||||
image_bytes = self.load_image(
|
||||
f"""{st.session_state.image_file_prefix}{st.session_state.image_to_edit + 1}.png"""
|
||||
)
|
||||
bg_image = Image.open(image_bytes)
|
||||
|
||||
st.markdown("<h1>Edit Image</h1>", unsafe_allow_html=True)
|
||||
|
||||
# Stroke Width Control
|
||||
# - Add a slider to control drawing/mask stroke width
|
||||
self.stroke_width = st.slider(
|
||||
"Stroke width: ",
|
||||
10,
|
||||
50,
|
||||
self.stroke_width,
|
||||
key="canvas_slider",
|
||||
)
|
||||
|
||||
# Image Prompt Section
|
||||
with st.form("Image prompt"):
|
||||
# - Provide a description of the form's purpose
|
||||
st.write("Input a query to generate the product.")
|
||||
img_prompt = st.text_input("Enter your custom query", "")
|
||||
edit_img_btn = st.form_submit_button("Submit prompt", type="primary")
|
||||
|
||||
# - Handle form submission
|
||||
if edit_img_btn:
|
||||
# -- Update session state to trigger image generation
|
||||
st.session_state.generate_images = True
|
||||
st.session_state.image_prompt = img_prompt
|
||||
|
||||
# Mask Drawing Setup
|
||||
drawing_dict = { # - Dictionary mapping descriptive names to drawing modes.
|
||||
"⬜ Rectangle": "rect",
|
||||
"🖌️ Brush": "freedraw",
|
||||
"⚪ Circle": "circle",
|
||||
"📏 Move/Scale/Rotate": "transform",
|
||||
}
|
||||
self.drawing_mode = st.selectbox(
|
||||
"[Optional] Draw a mask where you want to edit the image",
|
||||
drawing_dict.keys(),
|
||||
key="canvas_select_box",
|
||||
)
|
||||
|
||||
# Canvas Setup
|
||||
height = (
|
||||
int(bg_image.size[1] / (bg_image.size[0] / 704)) // 2
|
||||
) # - Calculate canvas height dynamically
|
||||
canvas_result = st_canvas(
|
||||
fill_color="rgba(255, 255, 255, 1)",
|
||||
stroke_width=self.stroke_width,
|
||||
stroke_color="rgba(255, 255, 255, 1)",
|
||||
background_color="#000",
|
||||
background_image=bg_image, # Use loaded image as background
|
||||
update_streamlit=self.realtime_update,
|
||||
height=height,
|
||||
initial_drawing=None,
|
||||
width=352,
|
||||
drawing_mode=drawing_dict[self.drawing_mode],
|
||||
point_display_radius=0, # - Hide cursor on canvas
|
||||
key="canvas",
|
||||
)
|
||||
|
||||
# Background Editing Control
|
||||
if st.checkbox("Edit Image Background"):
|
||||
st.session_state.bg_editing = True # - Enable background editing mode
|
||||
st.write(" Please mask the area you want to preserve")
|
||||
else:
|
||||
st.session_state.bg_editing = False # - Disable background editing
|
||||
|
||||
# Return Values
|
||||
return (
|
||||
canvas_result,
|
||||
bg_image,
|
||||
image_bytes,
|
||||
) # - Return values likely used elsewhere
|
||||
@@ -0,0 +1,33 @@
|
||||
"""
|
||||
Defines functions for generating text embeddings using a Vertex AI
|
||||
TextEmbeddingModel.
|
||||
"""
|
||||
|
||||
import backoff
|
||||
from google.api_core.exceptions import ResourceExhausted
|
||||
import numpy as np
|
||||
import streamlit as st
|
||||
from vertexai.preview.language_models import TextEmbeddingModel
|
||||
|
||||
|
||||
@st.cache_resource
|
||||
def get_embedding_model() -> TextEmbeddingModel:
|
||||
"""
|
||||
Loads embedding model (to be cached).
|
||||
"""
|
||||
embedding_model = TextEmbeddingModel.from_pretrained("text-embedding-005")
|
||||
return embedding_model
|
||||
|
||||
|
||||
@backoff.on_exception(backoff.expo, ResourceExhausted, max_time=10)
|
||||
def embedding_model_with_backoff(text: list[str]) -> np.ndarray:
|
||||
"""
|
||||
Process embeddings for uploaded files.
|
||||
Args:
|
||||
text: A list of text strings to process.
|
||||
|
||||
Returns:
|
||||
A NumPy array containing the processed embeddings.
|
||||
"""
|
||||
embeddings = get_embedding_model().get_embeddings(text)
|
||||
return np.array([each.values for each in embeddings][0])
|
||||
+184
@@ -0,0 +1,184 @@
|
||||
"""
|
||||
This module provides functions for creating content PDFs with specific
|
||||
layouts and formatting.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from app.pages_utils.pdf_generation import PDFRounded as pdf_generator
|
||||
from app.pages_utils.pdf_generation import add_formatted_page, check_add_page
|
||||
import streamlit as st
|
||||
|
||||
|
||||
def create_pdf_layout(
|
||||
pdf: pdf_generator, content: list[str], title: str, images: list[str]
|
||||
) -> None:
|
||||
"""
|
||||
Creates a PDF layout with the given content, title, and images.
|
||||
|
||||
Args:
|
||||
pdf: The PDF object where the layout will be created.
|
||||
content: A list of strings representing the textual content of the PDF.
|
||||
title: The title of the PDF.
|
||||
images: A list of image file names to include in the PDF.
|
||||
"""
|
||||
|
||||
for j, text in enumerate(content):
|
||||
add_formatted_page(pdf)
|
||||
|
||||
# Set up header
|
||||
pdf.set_xy(15, 15)
|
||||
pdf.set_text_color(106, 144, 226)
|
||||
pdf.set_font("Arial", "B", 11)
|
||||
pdf.multi_cell(
|
||||
180,
|
||||
5,
|
||||
f"{title} {st.session_state.product_category}",
|
||||
0,
|
||||
align="C",
|
||||
)
|
||||
|
||||
# Reset text color
|
||||
pdf.set_text_color(0, 0, 0)
|
||||
|
||||
# Add image
|
||||
pdf.set_font("Arial", "B", 11)
|
||||
pdf.set_xy(17, 25)
|
||||
image_path = f"gen_image{images[j]}.png"
|
||||
pdf.image(image_path, x=60, y=40, w=90, h=70)
|
||||
|
||||
# Add text content, handling potential page breaks
|
||||
pages = check_add_page(pdf, text)
|
||||
pdf.set_font("Arial", "", 11)
|
||||
|
||||
for i, page in enumerate(pages):
|
||||
if page.strip() == "":
|
||||
continue
|
||||
|
||||
if i == 0: # First page of text
|
||||
pdf.set_xy(17, 120)
|
||||
else: # Subsequent pages
|
||||
add_formatted_page(pdf)
|
||||
pdf.set_xy(17, 15)
|
||||
pdf.set_font("Arial", "", 11)
|
||||
|
||||
pdf.multi_cell(170, 5, page) # Output the text
|
||||
|
||||
|
||||
def create_content_pdf(
|
||||
product_content: list[list[dict[str, Any]]], selected_titles: list[str]
|
||||
) -> None:
|
||||
"""Creates a PDF for each product content and selected title.
|
||||
|
||||
Args:
|
||||
product_content: A list of strings representing the
|
||||
product content.
|
||||
selected_titles: A list of selected titles for each product content.
|
||||
"""
|
||||
for product_index in range(len(product_content) - 1):
|
||||
pdf = pdf_generator() # Create a PDF for the current product
|
||||
|
||||
# Build content and image lists for the current product
|
||||
content = [product_content[(int)(product_index)][0]["text"].replace("**", "")]
|
||||
images = [st.session_state.num_drafts * product_index + 1]
|
||||
|
||||
# Generate the PDF layout
|
||||
create_pdf_layout(pdf, content, selected_titles[product_index], images)
|
||||
|
||||
# Save the PDF with an appropriate filename
|
||||
pdf.output(f"content_{product_index}.pdf")
|
||||
|
||||
|
||||
def cut_string(string: str, num_characters: int) -> str:
|
||||
"""Cuts a string to the specified number of characters.
|
||||
|
||||
Args:
|
||||
string: The string to cut.
|
||||
num_characters: The number of characters to cut the string to.
|
||||
|
||||
Returns:
|
||||
The cut string.
|
||||
"""
|
||||
if len(string) <= num_characters:
|
||||
return string
|
||||
return string[:num_characters]
|
||||
|
||||
|
||||
def create_email_pdf(
|
||||
title: str, email_text: str, filename: str, image_name: str
|
||||
) -> None:
|
||||
"""Creates a PDF document from an email.
|
||||
|
||||
The PDF document contains the email subject, body, and an image.
|
||||
The title of the PDF document is set to the title of the email.
|
||||
|
||||
Args:
|
||||
title: The title of the email.
|
||||
email_text: The body of the email.
|
||||
filename: The name of the PDF file to be created.
|
||||
image_name: The name of the image file to be included in the PDF
|
||||
document.
|
||||
"""
|
||||
pdf = pdf_generator()
|
||||
|
||||
# Extract subject and text from email text.
|
||||
parts = email_text.split("\n", 1)
|
||||
subject = parts[0]
|
||||
text = parts[1]
|
||||
|
||||
# Add first page of pdf.
|
||||
add_formatted_page(pdf)
|
||||
|
||||
# Set location and text style for heading.
|
||||
pdf.set_xy(15, 15)
|
||||
pdf.set_text_color(106, 144, 226)
|
||||
pdf.set_font("Arial", "B", 11)
|
||||
|
||||
# Add heading to pdf object.
|
||||
pdf.multi_cell(
|
||||
180,
|
||||
5,
|
||||
f"{title} {st.session_state.product_category}",
|
||||
0,
|
||||
align="C",
|
||||
)
|
||||
|
||||
# Set text location and styling for subject.
|
||||
pdf.set_text_color(0, 0, 0)
|
||||
pdf.set_xy(17, 25)
|
||||
|
||||
# Add subject to pdf.
|
||||
pdf.multi_cell(180, 5, subject, 0, align="C")
|
||||
|
||||
# Add image to pdf object.
|
||||
pdf.image(f"{image_name}", x=60, y=40, w=90, h=70)
|
||||
|
||||
# Check if new page needs to be added, and
|
||||
# add required pages.
|
||||
# List pages stores the text content for each page.
|
||||
pages = check_add_page(pdf, text)
|
||||
|
||||
# Set font style for email body.
|
||||
pdf.set_font("Arial", "", 11)
|
||||
|
||||
# Add text content to each page of pdf.
|
||||
for i, page in enumerate(pages):
|
||||
# Check if an empty page is encountered.
|
||||
if page.strip() == "":
|
||||
continue
|
||||
|
||||
# First page
|
||||
if i == 0:
|
||||
pdf.set_xy(17, 120)
|
||||
pdf.multi_cell(170, 5, page)
|
||||
|
||||
# Remaining pages after the first page
|
||||
else:
|
||||
# Add new page.
|
||||
add_formatted_page(pdf)
|
||||
pdf.set_font("Arial", "", 11)
|
||||
|
||||
pdf.set_xy(17, 15)
|
||||
pdf.multi_cell(170, 5, page)
|
||||
|
||||
pdf.output(f"{filename}.pdf", "F")
|
||||
+75
@@ -0,0 +1,75 @@
|
||||
"""
|
||||
This module provides functions for interacting with Vertex AI text generation
|
||||
model (Gemini-Pro).
|
||||
|
||||
* generate_gemini():
|
||||
* Utilizes the Gemini-Pro model for flexible text generation.
|
||||
* Supports customization of generation parameters.
|
||||
* Incorporates safety settings.
|
||||
|
||||
* parallel_generate_search_results():
|
||||
* Employs asynchronous requests to Gemini for search result generation.
|
||||
* Handles potential errors during communication with the model.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import aiohttp as cloud_function_call
|
||||
from dotenv import load_dotenv
|
||||
import streamlit as st
|
||||
import vertexai
|
||||
from vertexai import generative_models
|
||||
|
||||
load_dotenv()
|
||||
|
||||
|
||||
PROJECT_ID = os.getenv("PROJECT_ID")
|
||||
LOCATION = os.getenv("LOCATION")
|
||||
|
||||
vertexai.init(project=PROJECT_ID, location=LOCATION)
|
||||
|
||||
|
||||
def generate_gemini(text_prompt: str) -> str:
|
||||
"""Generates text using the Gemini-Pro model.
|
||||
|
||||
Args:
|
||||
text_prompt: The text prompt to generate from.
|
||||
|
||||
Returns:
|
||||
The generated text.
|
||||
"""
|
||||
model = generative_models.GenerativeModel("gemini-2.0-flash")
|
||||
response = model.generate_content(
|
||||
text_prompt,
|
||||
generation_config=st.session_state.generation_config,
|
||||
)
|
||||
return response.text
|
||||
|
||||
|
||||
async def parallel_generate_search_results(query: str) -> str:
|
||||
"""Generates search results using the Gemini model in a parallel
|
||||
fashion.
|
||||
|
||||
Args:
|
||||
query: The query to generate search results for.
|
||||
|
||||
Returns:
|
||||
The generated search results.
|
||||
"""
|
||||
text_query = json.dumps({"text_prompt": query})
|
||||
async with cloud_function_call.ClientSession() as session:
|
||||
url = f"https://us-central1-{PROJECT_ID}.cloudfunctions.net/gemini-call"
|
||||
# Create post request to get text.
|
||||
async with session.post(
|
||||
url,
|
||||
data=text_query,
|
||||
headers=st.session_state.headers,
|
||||
verify_ssl=False,
|
||||
) as text_response:
|
||||
if text_response.status == 200:
|
||||
# If response is valid, return generated text.
|
||||
response = await text_response.json()
|
||||
response_text = response["generated_text"]
|
||||
return response_text
|
||||
return ""
|
||||
@@ -0,0 +1,119 @@
|
||||
"""
|
||||
Utility module to:
|
||||
- Resize image bytes
|
||||
- Generate an image with Imagen
|
||||
- Edit an image with Imagen
|
||||
- Render the image generation and editing UI
|
||||
"""
|
||||
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
|
||||
from PIL import Image
|
||||
import aiohttp as cloud_function_call
|
||||
import streamlit as st
|
||||
import vertexai
|
||||
from vertexai.preview.vision_models import ImageGenerationModel
|
||||
|
||||
logging.basicConfig(format="%(levelname)s:%(message)s", level=logging.DEBUG)
|
||||
|
||||
# Set project parameters
|
||||
PROJECT_ID = os.getenv("PROJECT_ID")
|
||||
LOCATION = os.getenv("LOCATION")
|
||||
|
||||
# Set project parameters
|
||||
IMAGE_MODEL_NAME = "imagegeneration@006"
|
||||
model = ImageGenerationModel.from_pretrained(IMAGE_MODEL_NAME)
|
||||
vertexai.init(project=PROJECT_ID, location=LOCATION)
|
||||
|
||||
|
||||
def predict_edit_image(
|
||||
instance_dict: dict,
|
||||
parameters: dict,
|
||||
) -> list[str]:
|
||||
"""Predicts the output of Imagen on a given instance dict.
|
||||
Args:
|
||||
instance_dict:
|
||||
The input to the large language model. (dict)
|
||||
parameters:
|
||||
The parameters for the prediction. (dict)
|
||||
Returns:
|
||||
A list of <vertexai.preview.vision_models.GeneratedImage> object
|
||||
containing the predictions.
|
||||
"""
|
||||
|
||||
responses = model.edit_image(
|
||||
prompt=instance_dict["prompt"],
|
||||
base_image=instance_dict["image"],
|
||||
# Optional parameters
|
||||
number_of_images=parameters["sampleCount"],
|
||||
language="en",
|
||||
mask=instance_dict["mask"],
|
||||
)
|
||||
return responses
|
||||
|
||||
|
||||
def image_generation(
|
||||
prompt: str,
|
||||
sample_count: int,
|
||||
aspect_ratio: str,
|
||||
filename: str,
|
||||
) -> None:
|
||||
"""Generates an image from a prompt.
|
||||
|
||||
Args:
|
||||
prompt:
|
||||
The prompt to use to generate the image.
|
||||
sample_count:
|
||||
The number of images to generate.
|
||||
aspect_ratio:
|
||||
The aspect ratio of the generated images.
|
||||
filename:
|
||||
The filename to store the image.
|
||||
|
||||
Returns:
|
||||
None.
|
||||
"""
|
||||
images = model.generate_images(
|
||||
prompt=prompt,
|
||||
# Optional parameters
|
||||
number_of_images=sample_count,
|
||||
language="en",
|
||||
aspect_ratio=aspect_ratio,
|
||||
)
|
||||
images[0].save(location=f"{filename}.png", include_generation_parameters=False)
|
||||
|
||||
|
||||
async def parallel_image_generation(prompt: str, col: int) -> Image.Image | None:
|
||||
"""
|
||||
Executes parallel generation of images through Imagen.
|
||||
|
||||
Args:
|
||||
prompt (String): Prompt for image Generation.
|
||||
col (int): A pointer to the draft number of the image.
|
||||
"""
|
||||
image_prompt = json.dumps({"img_prompt": prompt})
|
||||
async with cloud_function_call.ClientSession() as session:
|
||||
url = f"https://us-central1-{PROJECT_ID}.cloudfunctions.net/imagen-call"
|
||||
# Create a post request to get images.
|
||||
async with session.post(
|
||||
url,
|
||||
data=image_prompt,
|
||||
headers=st.session_state.headers,
|
||||
verify_ssl=False,
|
||||
) as img_response:
|
||||
# Check if response is valid.
|
||||
if img_response.status == 200:
|
||||
response = await img_response.read()
|
||||
|
||||
# Load response image.
|
||||
response_image = Image.open(io.BytesIO(response))
|
||||
# Save image for further use.
|
||||
response_image.save(
|
||||
f"gen_image{st.session_state.num_drafts+col}.png", format="PNG"
|
||||
)
|
||||
return response_image
|
||||
|
||||
return None
|
||||
@@ -0,0 +1,165 @@
|
||||
"""
|
||||
This module provides functions for generating insights and searching
|
||||
relevant information from uploaded data.
|
||||
This module:
|
||||
* Retrieves relevant context from a vector database based on a user's
|
||||
query.
|
||||
* Leverages Gemini-Pro to generate a precise answer.
|
||||
* Presents the answer along with top-matched context sources.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
|
||||
from app.pages_utils.embedding_model import embedding_model_with_backoff
|
||||
from app.pages_utils.get_llm_response import generate_gemini
|
||||
from app.pages_utils.pages_config import GLOBAL_CFG
|
||||
from dotenv import load_dotenv
|
||||
from google.cloud import storage
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import streamlit as st
|
||||
|
||||
load_dotenv()
|
||||
|
||||
|
||||
PROJECT_ID = os.getenv("PROJECT_ID")
|
||||
LOCATION = os.getenv("LOCATION")
|
||||
|
||||
# Define storage bucket
|
||||
storage_client = storage.Client(project=PROJECT_ID)
|
||||
bucket = storage_client.bucket(GLOBAL_CFG["bucket_name"])
|
||||
|
||||
|
||||
def extract_bullet_points(text: str) -> list[str]:
|
||||
"""
|
||||
Extracts all text enclosed within <b> and </b> tags
|
||||
or between ** tags from given string.
|
||||
|
||||
Args:
|
||||
html_string: The HTML string to process.
|
||||
|
||||
Returns:
|
||||
A list containing the extracted text segments.
|
||||
"""
|
||||
pattern = r"(?:<b>([^<]+?)</b>)|(?:\*\*(.+?)\*\*)"
|
||||
matches = re.findall(pattern, text)
|
||||
|
||||
# Flatten and filter out empty matches
|
||||
bold_text = [match for group in matches for match in group if match.strip()]
|
||||
|
||||
return bold_text
|
||||
|
||||
|
||||
def get_suggestions(state_key: str) -> None:
|
||||
"""Gets suggestions for the given state key.
|
||||
|
||||
Args:
|
||||
state_key (str): The state key to get suggestions for.
|
||||
"""
|
||||
|
||||
if st.session_state.rag_search_term is None:
|
||||
embeddings_df = st.session_state["processed_data_list"].head(2)
|
||||
context = "\n".join(embeddings_df["content"].values)
|
||||
prompt = f""" Context: \n {context} \n
|
||||
generate 5 questions based on the given context
|
||||
"""
|
||||
else:
|
||||
context = st.session_state.rag_search_term
|
||||
prompt = f""" Context: \n {context} \n
|
||||
generate 5 questions based on the given context. The questions
|
||||
should strictly be questions for further analysis of
|
||||
{st.session_state.rag_search_term}
|
||||
"""
|
||||
gen_suggestions = generate_gemini(prompt)
|
||||
st.session_state[state_key] = extract_bullet_points(gen_suggestions)
|
||||
|
||||
|
||||
def get_stored_embeddings_as_df() -> pd.DataFrame | None:
|
||||
"""Retrieves and processes stored embeddings from cloud storage.
|
||||
|
||||
Returns:
|
||||
A Pandas DataFrame containing the embeddings, or None if not found.
|
||||
"""
|
||||
embedding = bucket.blob(st.session_state.product_category + "/embeddings.json")
|
||||
|
||||
if embedding.exists():
|
||||
stored_embedding_data = embedding.download_as_string()
|
||||
embedding_dataframe = pd.DataFrame.from_dict(json.loads(stored_embedding_data))
|
||||
st.session_state["processed_data_list"] = embedding_dataframe
|
||||
return embedding_dataframe
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def get_filter_context_from_vector_database(
|
||||
question: str, sort_index_value: int = 3
|
||||
) -> tuple[str, pd.DataFrame]:
|
||||
"""Gets the filter context from the vector database.
|
||||
|
||||
Args:
|
||||
question (str): The question to get the filter context for.
|
||||
sort_index_value (int, optional): The number of top matched results
|
||||
to return.
|
||||
# Defaults to 3.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing the filter context and the top matched
|
||||
results.
|
||||
"""
|
||||
st.session_state["query_vectors"] = np.array(
|
||||
embedding_model_with_backoff([question])
|
||||
)
|
||||
top_matched_score = (
|
||||
st.session_state["processed_data_list"]["embedding"]
|
||||
.apply(
|
||||
lambda row: (
|
||||
np.dot(row, st.session_state["query_vectors"]) if row is not None else 0
|
||||
)
|
||||
)
|
||||
.sort_values(ascending=False)[:sort_index_value]
|
||||
)
|
||||
|
||||
top_matched_df = st.session_state["processed_data_list"][
|
||||
st.session_state["processed_data_list"].index.isin(top_matched_score.index)
|
||||
]
|
||||
top_matched_df = top_matched_df[["file_name", "chunk_number", "content"]]
|
||||
top_matched_df["confidence_score"] = top_matched_score
|
||||
top_matched_df.sort_values(by=["confidence_score"], ascending=False, inplace=True)
|
||||
|
||||
context = "\n".join(
|
||||
st.session_state["processed_data_list"][
|
||||
st.session_state["processed_data_list"].index.isin(top_matched_score.index)
|
||||
]["content"].values
|
||||
)
|
||||
return (context, top_matched_df)
|
||||
|
||||
|
||||
def generate_insights_search_result(query: str) -> tuple[str, pd.DataFrame]:
|
||||
"""Generates insights search results for the given query.
|
||||
|
||||
Args:
|
||||
query (str): The query to generate insights search results for.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing the insights answer and the top matched
|
||||
results.
|
||||
"""
|
||||
|
||||
question = query
|
||||
context, top_matched_df = get_filter_context_from_vector_database(
|
||||
question=query, sort_index_value=20
|
||||
)
|
||||
|
||||
question_prompt_template = f"""
|
||||
Answer the question as precise as possible using the provided context.
|
||||
If the answer is not contained in the context, say "answer not available
|
||||
in context"
|
||||
\n \n
|
||||
Context: \n {context} \n
|
||||
Question: \n {question} \n
|
||||
Answer:"""
|
||||
|
||||
insights_answer = generate_gemini(question_prompt_template)
|
||||
return insights_answer, top_matched_df.head(5)
|
||||
@@ -0,0 +1,56 @@
|
||||
"""
|
||||
Utility module to work with app config.
|
||||
"""
|
||||
|
||||
from os.path import isfile
|
||||
|
||||
import pytomlpp as config_loader
|
||||
|
||||
APP_TOML = "./app/app_config.toml"
|
||||
OVERRIDE_TOML = "./override.toml"
|
||||
|
||||
assert isfile(APP_TOML), f"The file {APP_TOML} should exist"
|
||||
|
||||
with open(APP_TOML, "rb") as f:
|
||||
try:
|
||||
data = config_loader.load(f)
|
||||
except config_loader.DecodeError as e:
|
||||
print("Invalid App Configuration TOML file.")
|
||||
print(str(e))
|
||||
raise
|
||||
|
||||
|
||||
def merge(a: dict, b: dict) -> None:
|
||||
"""
|
||||
merge dictionaries a and b.
|
||||
"""
|
||||
for key in b:
|
||||
if key in a:
|
||||
if isinstance(a[key], dict) and isinstance(b[key], dict):
|
||||
merge(a[key], b[key])
|
||||
elif a[key] != b[key]:
|
||||
a[key] = b[key]
|
||||
else:
|
||||
a[key] = b[key]
|
||||
|
||||
|
||||
if isfile(OVERRIDE_TOML):
|
||||
with open(OVERRIDE_TOML, "rb") as f:
|
||||
try:
|
||||
data_override = config_loader.load(f)
|
||||
merge(data, data_override)
|
||||
except config_loader.DecodeError as e:
|
||||
print("Invalid Override TOML File")
|
||||
print(str(e))
|
||||
except Exception as e:
|
||||
print("Unexpected error")
|
||||
print(str(e))
|
||||
raise
|
||||
|
||||
assert "translate_api" in data, "No translation options in the config"
|
||||
assert "pages" in data, "No page configurations found in the config"
|
||||
|
||||
|
||||
TRANSLATE_CFG = data["translate_api"]
|
||||
PAGES_CFG = data["pages"]
|
||||
GLOBAL_CFG = data["global"]
|
||||
@@ -0,0 +1,227 @@
|
||||
"""
|
||||
This module provides functions for adding and formatting content within PDF
|
||||
documents.
|
||||
|
||||
* add_formatted_page(pdf):
|
||||
* Adds a standard-format page with a light gray background and a centered
|
||||
white rectangle.
|
||||
|
||||
* check_add_page(pdf, text):
|
||||
* Handles potential text overflow onto subsequent pages.
|
||||
|
||||
* Provides class for generating a pdf template for exporting content and
|
||||
emails.
|
||||
"""
|
||||
|
||||
# pylint: disable=R0913
|
||||
|
||||
from math import sqrt
|
||||
|
||||
import fpdf as pdf_generator
|
||||
|
||||
|
||||
class PDFRounded(pdf_generator.FPDF):
|
||||
"""
|
||||
Initializes basic PDF template for email and content files
|
||||
"""
|
||||
|
||||
def rounded_rect(
|
||||
self,
|
||||
x: float,
|
||||
y: float,
|
||||
w: float,
|
||||
h: float,
|
||||
r: float,
|
||||
style: str = "",
|
||||
corners: str = "1234",
|
||||
) -> None:
|
||||
"""
|
||||
Draws a rectangle with rounded corners.
|
||||
|
||||
Args:
|
||||
x (float): The x-coordinate of the top-left corner of the
|
||||
rectangle.
|
||||
y (float): The y-coordinate of the top-left corner of the
|
||||
rectangle.
|
||||
w (float): The width of the rectangle.
|
||||
h (float): The height of the rectangle.
|
||||
r (float): The radius of the rounded corners.
|
||||
style (str, optional): The style of the rectangle.
|
||||
Can be 'F' for filled, 'FD' for filled and drawn,
|
||||
or 'DF' for drawn and filled. Defaults to 'S' for stroked.
|
||||
corners (str, optional): A string of characters indicating
|
||||
which corners of the rectangle should be rounded.
|
||||
Can be '1234' for all corners, '12' for the top-left and top-right
|
||||
corners,
|
||||
'34' for the bottom-left and bottom-right corners, or
|
||||
'13' for the top-left and bottom-right corners. Defaults to '1234'.
|
||||
"""
|
||||
|
||||
k = self.k
|
||||
hp = self.h
|
||||
if style == "F":
|
||||
op = "f"
|
||||
elif style in ("FD", "DF"):
|
||||
op = "B"
|
||||
else:
|
||||
op = "S"
|
||||
my_arc = 4 / 3 * (sqrt(2) - 1)
|
||||
self._out(f"{(x + r) * k} {(hp - y) * k} m")
|
||||
|
||||
xc = x + w - r
|
||||
yc = y + r
|
||||
self._out(f"{xc * k} {(hp - y) * k} l")
|
||||
if "2" not in corners:
|
||||
self._out(f"{(x + w) * k} {(hp - y) * k} l")
|
||||
else:
|
||||
self.arc(
|
||||
xc + r * my_arc,
|
||||
yc - r,
|
||||
xc + r,
|
||||
yc - r * my_arc,
|
||||
xc + r,
|
||||
yc,
|
||||
)
|
||||
|
||||
xc = x + w - r
|
||||
yc = y + h - r
|
||||
self._out(f"{(x + w) * k} {(hp - yc) * k} l")
|
||||
if "3" not in corners:
|
||||
self._out(f"{(x + w) * k} {(hp - (y + h)) * k} l")
|
||||
else:
|
||||
self.arc(
|
||||
xc + r,
|
||||
yc + r * my_arc,
|
||||
xc + r * my_arc,
|
||||
yc + r,
|
||||
xc,
|
||||
yc + r,
|
||||
)
|
||||
|
||||
xc = x + r
|
||||
yc = y + h - r
|
||||
self._out(f"{xc * k} {(hp - (y + h)) * k} l")
|
||||
if "4" not in corners:
|
||||
self._out(f"{x * k} {(hp - (y + h)) * k} l")
|
||||
else:
|
||||
self.arc(
|
||||
xc - r * my_arc,
|
||||
yc + r,
|
||||
xc - r,
|
||||
yc + r * my_arc,
|
||||
xc - r,
|
||||
yc,
|
||||
)
|
||||
|
||||
xc = x + r
|
||||
yc = y + r
|
||||
self._out(f"{x * k} {(hp - yc) * k} l")
|
||||
if "1" not in corners:
|
||||
self._out(f"{x * k} {(hp - y) * k} l")
|
||||
self._out(f"{(x + r) * k} {(hp - y) * k} l")
|
||||
else:
|
||||
self.arc(
|
||||
xc - r,
|
||||
yc - r * my_arc,
|
||||
xc - r * my_arc,
|
||||
yc - r,
|
||||
xc,
|
||||
yc - r,
|
||||
)
|
||||
self._out(op)
|
||||
|
||||
def arc(
|
||||
self, x1: float, y1: float, x2: float, y2: float, x3: float, y3: float
|
||||
) -> None:
|
||||
"""
|
||||
Draws an arc.
|
||||
|
||||
Args:
|
||||
x1 (float): The x-coordinate of the start point of the arc.
|
||||
y1 (float): The y-coordinate of the start point of the arc.
|
||||
x2 (float): The x-coordinate of the end point of the arc.
|
||||
y2 (float): The y-coordinate of the end point of the arc.
|
||||
x3 (float): The x-coordinate of the control point of the arc.
|
||||
y3 (float): The y-coordinate of the control point of the arc.
|
||||
"""
|
||||
h = self.h
|
||||
self._out(
|
||||
f"""{x1 * self.k:.2f}
|
||||
{(h - y1) * self.k:.2f} {x2 * self.k:.2f}
|
||||
{(h - y2) * self.k:.2f} {x3 * self.k:.2f}
|
||||
{(h - y3) * self.k:.2f} c"""
|
||||
)
|
||||
|
||||
|
||||
def add_formatted_page(pdf: pdf_generator) -> None:
|
||||
"""Adds a formatted page to the PDF document.
|
||||
|
||||
The page is filled with a light gray color and has a white rectangle in
|
||||
the center.
|
||||
The font is set to Arial, bold, size 18, and the fill color is set to
|
||||
white.
|
||||
|
||||
Args:
|
||||
pdf: The PDF document to which the page is added.
|
||||
"""
|
||||
pdf.set_fill_color(225, 230, 237)
|
||||
pdf.add_page()
|
||||
pdf.rect(0, 0, 210, 297, "F")
|
||||
pdf.set_font("Arial", "B", 18)
|
||||
pdf.set_fill_color(255, 255, 255)
|
||||
pdf.rect(10, 10, 190, 277, "F")
|
||||
|
||||
|
||||
def check_add_page(pdf: pdf_generator, text: str) -> list[str]:
|
||||
"""Checks if the text overflows onto a new page and adds a new page if
|
||||
necessary.
|
||||
|
||||
The text is split into lines based on the available page width.
|
||||
If the text overflows onto a new page, a new page is added and the text
|
||||
is continued on the new page.
|
||||
|
||||
Args:
|
||||
pdf: The PDF document to which the text is added.
|
||||
text: The text to be added to the PDF document.
|
||||
|
||||
Returns:
|
||||
A list of strings, where each page is a string containing the text
|
||||
that fits on that page.
|
||||
"""
|
||||
|
||||
pages: list[str] = [] # Store the text content for each page
|
||||
page_content = "" # Text for the current page being built
|
||||
|
||||
pdf.set_font("Arial", "", 11)
|
||||
|
||||
# Split the text into lines based on the available page width
|
||||
lines = []
|
||||
y = pdf.y
|
||||
|
||||
for line in text.split("\n"):
|
||||
words = line.split(" ")
|
||||
current_line = "" # Represents a single line within the page
|
||||
|
||||
for word in words:
|
||||
# Check if adding the word would exceed the line limit
|
||||
if len(current_line + word) > pdf.w - pdf.l_margin - pdf.r_margin:
|
||||
lines.append(current_line)
|
||||
page_content += current_line + "\n"
|
||||
current_line = "" # Reset for the next line
|
||||
current_line += word + " "
|
||||
|
||||
if current_line:
|
||||
lines.append(current_line)
|
||||
# Check if adding the line would overflow the page
|
||||
if y + 10 > pdf.h - (80 if len(pages) == 0 else 0):
|
||||
pages.append(page_content) # Add the new page to the PDF
|
||||
page_content = current_line + "\n" # Start a new page
|
||||
y = 10
|
||||
else:
|
||||
# Add the completed line to the current page
|
||||
page_content += current_line + "\n"
|
||||
y += 10
|
||||
|
||||
# Add any remaining text to the final page
|
||||
pages.append(page_content)
|
||||
return pages
|
||||
+215
@@ -0,0 +1,215 @@
|
||||
"""
|
||||
This module provides functions for rendering product feature drafts and
|
||||
managing user selections.
|
||||
This module:
|
||||
* Fetches a product feature response from the LLM, ensuring a specific
|
||||
format.
|
||||
* Displays draft features in a grid layout with checkboxes for selection.
|
||||
* Facilitates the modification of selected features, updating the UI
|
||||
accordingly.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from app.pages_utils.get_llm_response import generate_gemini
|
||||
from dotenv import load_dotenv
|
||||
import streamlit as st
|
||||
|
||||
logging.basicConfig(format="%(levelname)s:%(message)s", level=logging.DEBUG)
|
||||
|
||||
load_dotenv()
|
||||
|
||||
BOX_STYLE = """
|
||||
border: 0.5px solid #6a90e2;
|
||||
padding: 10px;
|
||||
margin: 10px;
|
||||
height: 280px;
|
||||
border-radius: 25px;
|
||||
"""
|
||||
|
||||
CLICKED_BOX_STYLE = """
|
||||
border: 2px solid #3367D6;
|
||||
padding: 10px;
|
||||
margin: 10px;
|
||||
height: 280px;
|
||||
border-radius: 25px;
|
||||
"""
|
||||
|
||||
|
||||
def _add_title_to_selection(title: str) -> None:
|
||||
"""Adds a title to the list of selected titles.
|
||||
|
||||
Args:
|
||||
title (str): The title to add.
|
||||
"""
|
||||
if title not in st.session_state.selected_titles:
|
||||
st.session_state.selected_titles.append(title)
|
||||
if st.session_state.modifying:
|
||||
st.rerun()
|
||||
|
||||
|
||||
def _remove_title_from_selection(title: str) -> None:
|
||||
"""Removes a title from the list of selected titles.
|
||||
|
||||
Args:
|
||||
title (str): The title to remove.
|
||||
"""
|
||||
if title in st.session_state.selected_titles:
|
||||
st.session_state.selected_titles.remove(title)
|
||||
|
||||
|
||||
def _render_box(box_id: str, title: str, parts: list, class_name: str) -> None:
|
||||
"""Renders a box with the given title, parts, and style.
|
||||
|
||||
Args:
|
||||
box_id (str): The ID of the box.
|
||||
title (str): The title of the box.
|
||||
parts (list): The parts of the box.
|
||||
style (str): The style of the box.
|
||||
"""
|
||||
st.markdown(
|
||||
f"""<div id={box_id} class={class_name}>
|
||||
<h5 style="color: #3367D6; text-align: center">
|
||||
{title if len(parts) == 2 else parts[0]}
|
||||
</h5>
|
||||
<div> {'' if len(parts) == 0 else parts[1]} </div>
|
||||
</div>
|
||||
""",
|
||||
unsafe_allow_html=True,
|
||||
)
|
||||
|
||||
|
||||
def get_features(text: str) -> list[str]:
|
||||
"""Gets a list of features from the given text.
|
||||
App displays a grid of boxes where each box corresponds to
|
||||
a particular feature. This function divides the given input
|
||||
text to a list of those features
|
||||
|
||||
Args:
|
||||
text (str): The text to get the points from.
|
||||
|
||||
Returns:
|
||||
list: A list of points.
|
||||
"""
|
||||
points = text.split("\n")
|
||||
curr = ""
|
||||
sep_points = []
|
||||
for point in points:
|
||||
point = point.strip()
|
||||
if point == "":
|
||||
continue
|
||||
if point.endswith(":"):
|
||||
curr = point
|
||||
else:
|
||||
if point.endswith("."):
|
||||
sep_points.append(curr.strip() + point.strip())
|
||||
curr = ""
|
||||
else:
|
||||
curr += point
|
||||
return sep_points
|
||||
|
||||
|
||||
def generate_formatted_response(prompt: str) -> str:
|
||||
"""Generates a formatted response based on the given prompt.
|
||||
|
||||
Args:
|
||||
prompt (str): The user-selected or custom prompt.
|
||||
|
||||
Returns:
|
||||
str: The formatted response text.
|
||||
"""
|
||||
with st.spinner("Fetching Response..."):
|
||||
generated_response = generate_gemini(
|
||||
f""" {prompt} in 12 points. The answer should strictly
|
||||
be a numbered list. Every bullet point should be strictly
|
||||
less than 150 characters. Each point should strictly have
|
||||
the format of title followed by description separated by ':'.
|
||||
Every bullet point should strictly have exactly one title.
|
||||
Use of bold text should be avoided strictly.
|
||||
Each title should strictly be a combination of 2 or
|
||||
more features."""
|
||||
)
|
||||
formatted_response = generated_response.replace("**", "")
|
||||
logging.debug(formatted_response)
|
||||
return formatted_response
|
||||
|
||||
|
||||
def render_features(features: st.delta_generator.DeltaGenerator) -> None:
|
||||
"""Renders draft ideas in a grid format, allowing for selection.
|
||||
|
||||
Args:
|
||||
features: A list of draft ideas to display.
|
||||
"""
|
||||
|
||||
if st.session_state.generated_points is None:
|
||||
st.session_state.generated_points = get_features(
|
||||
st.session_state.generated_response
|
||||
)
|
||||
|
||||
with features:
|
||||
col1, col2, col3 = st.columns(3)
|
||||
|
||||
for i, point in enumerate(st.session_state.generated_points):
|
||||
box = (
|
||||
col1 if i % 3 == 0 else col2 if i % 3 == 1 else col3
|
||||
) # Inline conditionals
|
||||
|
||||
box_id = f"box_{i}"
|
||||
|
||||
# Split point into two parts based on a colon (':') delimiter
|
||||
parts = point.split(":", 1) # Split at most once
|
||||
if len(parts) == 2:
|
||||
# Extract and clean up the parts
|
||||
first_part = parts[0].strip()
|
||||
second_part = parts[1].strip()
|
||||
parts = [first_part, second_part]
|
||||
# No colon found, part assigned as the whole sentence
|
||||
else:
|
||||
parts = [point]
|
||||
|
||||
# Trim title to only the heading
|
||||
title = parts[0]
|
||||
try:
|
||||
title_parts = title.split(".")
|
||||
title = title_parts[1].strip()
|
||||
except IndexError:
|
||||
logging.debug("Unable to trim title")
|
||||
|
||||
with box:
|
||||
checkbox_key = f"{title} {i}"
|
||||
|
||||
# Checkbox logic for idea selection
|
||||
checkbox = st.checkbox(
|
||||
"Select Idea",
|
||||
key=checkbox_key,
|
||||
value=title in st.session_state.selected_titles,
|
||||
)
|
||||
if checkbox:
|
||||
_add_title_to_selection(title)
|
||||
else:
|
||||
_remove_title_from_selection(title)
|
||||
|
||||
# Rendering with appropriate styles
|
||||
if title in st.session_state.selected_titles:
|
||||
_render_box(box_id, title, parts, "box-clicked")
|
||||
else:
|
||||
_render_box(box_id, title, parts, "box-default")
|
||||
|
||||
|
||||
def modify_selection(content: st.container) -> None:
|
||||
"""Modifies the selection of features.
|
||||
|
||||
Args:
|
||||
content: The streamlit container widget to modify.
|
||||
"""
|
||||
st.session_state.modifying = True
|
||||
new_features = st.empty()
|
||||
render_features(new_features)
|
||||
if st.session_state.content_generated is True:
|
||||
content.empty()
|
||||
content = st.empty()
|
||||
st.session_state.content_generated = False
|
||||
st.session_state.product_content = []
|
||||
st.session_state.create_product = False
|
||||
st.session_state.generate_images = False
|
||||
st.rerun()
|
||||
@@ -0,0 +1,281 @@
|
||||
"""
|
||||
This module provides functions for generating and managing product
|
||||
content based on selected features.
|
||||
Functions include:
|
||||
|
||||
* Initiate text and image generation with user-provided features.
|
||||
* Store the generated content for display.
|
||||
* Support content generation with asynchronous calls.
|
||||
* Render a form for selecting pre-defined prompts or entering custom
|
||||
queries.
|
||||
* Facilitate the generation of product feature suggestions.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from app.pages_utils.get_llm_response import (
|
||||
generate_gemini,
|
||||
parallel_generate_search_results,
|
||||
)
|
||||
from app.pages_utils.imagen import parallel_image_generation
|
||||
from dotenv import load_dotenv
|
||||
import streamlit as st
|
||||
|
||||
logging.basicConfig(
|
||||
format="%(levelname)s:%(message)s",
|
||||
level=logging.DEBUG,
|
||||
)
|
||||
load_dotenv()
|
||||
|
||||
|
||||
def update_generation_state() -> None:
|
||||
"""Updates the generation state post generate button click."""
|
||||
|
||||
# Check whether custom prompt has been given.
|
||||
if st.session_state.custom_prompt != "":
|
||||
st.session_state.selected_prompt = st.session_state.custom_prompt
|
||||
st.session_state.custom_prompt = ""
|
||||
|
||||
st.session_state.features_generated = True # Initiate feature generation
|
||||
st.session_state.generated_response = (
|
||||
None # store the response by llm for features.
|
||||
)
|
||||
# Track whether content corresponding to features has been generated.
|
||||
st.session_state.content_generated = False
|
||||
st.session_state.create_product = (
|
||||
False # Tracks whether a new product idea has been created.
|
||||
)
|
||||
st.session_state.selected_titles = (
|
||||
[]
|
||||
) # Stores selected titles for new product generation.
|
||||
st.session_state.product_content = [] # Content corresponding to each feature.
|
||||
st.session_state.content_edited = False # Tracks whether content is being edited.
|
||||
|
||||
|
||||
def generate_product_suggestions_for_feature_generation() -> None:
|
||||
"""Generates suggestions for a given product category for feature
|
||||
generation.
|
||||
|
||||
Returns:
|
||||
list: A list of feature suggestions.
|
||||
"""
|
||||
with st.spinner("Fetching Suggestions..."):
|
||||
feature_prompts = generate_gemini(
|
||||
f"""5 broad categories of {st.session_state.product_category}
|
||||
buyers. Give answer as a numbered list. Each point should
|
||||
strictly be only a category without any description."""
|
||||
)
|
||||
st.session_state.feature_suggestions = create_suggestion_list(feature_prompts)
|
||||
|
||||
|
||||
def build_prompt_form() -> bool:
|
||||
"""Creates the form for selecting prompts and entering custom queries.
|
||||
|
||||
Returns:
|
||||
Boolean value indicating whether the form was submitted.
|
||||
"""
|
||||
if st.session_state.feature_suggestions is None:
|
||||
st.session_state.feature_suggestions = []
|
||||
generate_product_suggestions_for_feature_generation()
|
||||
with st.form("prompt input"):
|
||||
options = [
|
||||
f"""Recommend {st.session_state.product_category} formulation
|
||||
features for {segment}"""
|
||||
for segment in st.session_state.feature_suggestions
|
||||
if st.session_state.feature_suggestions is not None
|
||||
]
|
||||
st.session_state.selected_prompt = st.selectbox(
|
||||
"Select an option or enter a custom query",
|
||||
options,
|
||||
)
|
||||
st.session_state.custom_prompt = st.text_input(
|
||||
"Enter your custom query",
|
||||
st.session_state.custom_prompt,
|
||||
)
|
||||
st.session_state.num_drafts = 1
|
||||
return st.form_submit_button("Generate", type="primary")
|
||||
|
||||
|
||||
def create_suggestion_list(gen_suggestions: str) -> list[str]:
|
||||
"""Creates a list of suggestions from the generated suggestions.
|
||||
|
||||
Args:
|
||||
gen_suggestions (str): The generated suggestions.
|
||||
|
||||
Returns:
|
||||
list: A list of suggestions.
|
||||
"""
|
||||
suggestions = []
|
||||
sep_suggestions = gen_suggestions.split("\n")
|
||||
for suggestion in sep_suggestions:
|
||||
suggestion_split = suggestion.split(".")
|
||||
if len(suggestion_split) > 1:
|
||||
suggestions.append(suggestion.split(".", 1)[1])
|
||||
|
||||
return suggestions
|
||||
|
||||
|
||||
async def parallel_call(titles: list[str]) -> list[Any]:
|
||||
"""
|
||||
Performs parallel calls to the text and image generation APIs.
|
||||
|
||||
Args:
|
||||
titles (list): A list of product titles.
|
||||
|
||||
Returns:
|
||||
list: A list of tuples containing the text and image generation
|
||||
results.
|
||||
"""
|
||||
logging.debug("entered parallel call")
|
||||
text_processes = []
|
||||
img_processes = []
|
||||
for index, title in enumerate(titles):
|
||||
# Handle edge case (No assorted products to be created in case only
|
||||
# one feature is selected)
|
||||
if index == 1 and len(st.session_state.selected_titles) == 1:
|
||||
break
|
||||
|
||||
# Create image generation and text generation prompts.
|
||||
img_prompt = f"{st.session_state.product_category} with {title} packaging."
|
||||
text_prompt = f"""Generate an innovative and original idea for a
|
||||
{st.session_state.product_category} that is {title} for
|
||||
{st.session_state.selected_prompt}. List ingredients of the suggested
|
||||
product. List benefits for different demographics of consumers of the
|
||||
product. The answer should strictly be very long and detailed and
|
||||
capture all features of the suggested product. Separately give the
|
||||
utility of the product for any three example consumer segments.
|
||||
Strictly demonstrate how the suggested product is an improvement
|
||||
over existing products."""
|
||||
|
||||
# Parallel calls to generate new content.
|
||||
if st.session_state.content_generated is False:
|
||||
text_processes.append(
|
||||
asyncio.create_task(parallel_generate_search_results(text_prompt))
|
||||
)
|
||||
img_processes.append(
|
||||
asyncio.create_task(parallel_image_generation(img_prompt, index))
|
||||
)
|
||||
|
||||
# Append the generated content to final result arrays.
|
||||
text_result_arr = await asyncio.gather(*text_processes)
|
||||
image_result_arr = await asyncio.gather(*img_processes)
|
||||
|
||||
return [text_result_arr, image_result_arr]
|
||||
|
||||
|
||||
async def prepare_titles() -> list[str]:
|
||||
"""Processes selected titles, handling edge cases.
|
||||
|
||||
Returns:
|
||||
list: A list of processed titles.
|
||||
"""
|
||||
titles = st.session_state.selected_titles.copy()
|
||||
|
||||
# Assorted titles to be created only if the length of selected features
|
||||
# is greater than 1.
|
||||
if len(st.session_state.selected_titles) > 1:
|
||||
# Create assorted product title if end of array is reached.
|
||||
# If end of array selected_titles array is not reached, keep original
|
||||
# title.
|
||||
assorted_title = ", ".join(st.session_state.selected_titles)
|
||||
|
||||
titles.append(assorted_title)
|
||||
# Store assorted title in session state.
|
||||
st.session_state.assorted_prod_title = assorted_title
|
||||
|
||||
return titles
|
||||
|
||||
|
||||
async def generate_product_content() -> None:
|
||||
"""
|
||||
Generates product content based on the selected titles and prompts.
|
||||
"""
|
||||
|
||||
if st.session_state.product_content is None:
|
||||
st.session_state.product_content = [] # Initialize product content storage
|
||||
|
||||
elements: list[list[dict[str, Any]]] = []
|
||||
|
||||
with st.spinner("Generating Product Ideas.."):
|
||||
# Fetch appropriate titles for processing
|
||||
titles = await prepare_titles()
|
||||
|
||||
# Call image and text generation function in parallel for efficiency
|
||||
task1 = asyncio.create_task(parallel_call(titles))
|
||||
result_array = await task1
|
||||
text_result_arr = result_array[0]
|
||||
|
||||
# Iterate over selected titles to generate content
|
||||
i = 0
|
||||
while i <= len(st.session_state.selected_titles):
|
||||
if i == 1 and len(st.session_state.selected_titles) == 1:
|
||||
break
|
||||
|
||||
# Prepare containers for the current iteration's content
|
||||
current_content = []
|
||||
|
||||
# Representing elements as a list of lists to handle multiple
|
||||
# drafts for same feature.
|
||||
elements.append([])
|
||||
|
||||
if i < len(st.session_state.selected_titles):
|
||||
title = titles[i]
|
||||
|
||||
# Generate content only if not already generated
|
||||
if st.session_state.content_generated is False:
|
||||
if i < len(st.session_state.selected_titles):
|
||||
current_content.append(text_result_arr[i])
|
||||
st.session_state.product_content.append(current_content)
|
||||
else:
|
||||
st.session_state.assorted_prod_content.append(text_result_arr[i])
|
||||
|
||||
# Build data for display elements
|
||||
elements[i].append(
|
||||
{
|
||||
"title": f"{title.strip()}",
|
||||
"text": (
|
||||
st.session_state.product_content[i][0].strip()
|
||||
if i < len(st.session_state.product_content)
|
||||
else st.session_state.assorted_prod_content[0]
|
||||
),
|
||||
"interval": None,
|
||||
"img": f"gen_image{st.session_state.num_drafts*i+1}.png",
|
||||
}
|
||||
)
|
||||
i += 1
|
||||
|
||||
# Store elements for display purposes
|
||||
st.session_state.draft_elements = elements
|
||||
|
||||
|
||||
async def handle_content_generation(features: st.container) -> None:
|
||||
"""
|
||||
Encapsulates the core content generation process.
|
||||
|
||||
Args:
|
||||
features (streamlit.container): A container to be cleared after
|
||||
content generation.
|
||||
"""
|
||||
|
||||
if not st.session_state.selected_titles:
|
||||
st.error("Please Select at least one Draft for Content Generation")
|
||||
return # Stop execution if no titles are selected
|
||||
|
||||
# features' is a UI element to be cleared
|
||||
features.empty()
|
||||
|
||||
await generate_product_content() # generates content
|
||||
|
||||
st.session_state.create_product = (
|
||||
True # Tracks whether product ideas have been generated.
|
||||
)
|
||||
st.session_state.content_generated = (
|
||||
True # Tracks whether product content has been generated.
|
||||
)
|
||||
|
||||
# Prepare titles for processing.
|
||||
st.session_state.chosen_titles = st.session_state.selected_titles.copy()
|
||||
if len(st.session_state.selected_titles) > 1:
|
||||
st.session_state.chosen_titles.append(st.session_state.assorted_prod_title)
|
||||
@@ -0,0 +1,126 @@
|
||||
"""
|
||||
This module provides functions for interacting with the Google Cloud Storage
|
||||
bucket, specifically for managing projects and their associated files.
|
||||
This module:
|
||||
* Retrieves a list of existing projects from the GCS bucket.
|
||||
* Updates the project list stored in the GCS bucket.
|
||||
* Lists PDF, text, and other supported file types in the current
|
||||
project's GCS bucket.
|
||||
* Deletes an entire project and its contents from the GCS bucket.
|
||||
* Deletes a specific file from the GCS project.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from app.pages_utils.pages_config import GLOBAL_CFG
|
||||
from dotenv import load_dotenv
|
||||
from google.cloud import storage
|
||||
import pandas as pd
|
||||
import streamlit as st
|
||||
|
||||
load_dotenv()
|
||||
|
||||
logging.basicConfig(format="%(levelname)s:%(message)s", level=logging.DEBUG)
|
||||
|
||||
PROJECT_ID = os.getenv("PROJECT_ID")
|
||||
LOCATION = os.getenv("LOCATION")
|
||||
|
||||
# Define storage bucket
|
||||
storage_client = storage.Client(project=PROJECT_ID)
|
||||
bucket = storage_client.bucket(GLOBAL_CFG["bucket_name"])
|
||||
|
||||
|
||||
def list_pdf_files_gcs() -> list[list[Any]]:
|
||||
"""Lists the PDF files in the current project's GCS bucket.
|
||||
|
||||
This function lists the PDF files in the current project's GCS bucket.
|
||||
It uses the 'storage_client' to get the list of blobs in the bucket and
|
||||
then filters the list to only include PDF files.
|
||||
It then returns a list of tuples of the blob name and the file extension.
|
||||
|
||||
Returns:
|
||||
list[list[Any]]: A list of tuples of the blob name and the file
|
||||
extension.
|
||||
"""
|
||||
project_embedding = bucket.blob(
|
||||
f"{st.session_state.product_category}/embeddings.json"
|
||||
)
|
||||
files = []
|
||||
if project_embedding.exists():
|
||||
file_list = bucket.list_blobs(prefix=f"{st.session_state.product_category}/")
|
||||
for file in file_list:
|
||||
_, file_extension = os.path.splitext(file.name)
|
||||
if file_extension in (".pdf", ".txt", ".csv", ".docx"):
|
||||
files.append([file.name, file_extension])
|
||||
else:
|
||||
st.write("No file uploaded")
|
||||
return files
|
||||
|
||||
|
||||
def delete_project_from_gcs() -> None:
|
||||
"""Deletes the current project from the GCS bucket.
|
||||
|
||||
This function deletes the current project from the GCS bucket.
|
||||
It uses the 'storage_client' to get the list of blobs in the bucket
|
||||
and then deletes all of the blobs in the bucket.
|
||||
It then removes the current project from the list of projects and
|
||||
updates the 'project_list.txt' file in the GCS bucket.
|
||||
"""
|
||||
# Load list of files for current project.
|
||||
project_file_list = bucket.list_blobs(
|
||||
prefix=f"{st.session_state.product_category}/"
|
||||
)
|
||||
|
||||
# Delete the files in the project.
|
||||
for file in project_file_list:
|
||||
file.delete()
|
||||
|
||||
# Remove the project name corresponding to the deleted project.
|
||||
st.session_state.product_categories.remove(st.session_state.product_category)
|
||||
|
||||
# Reset selected project to next project in list.
|
||||
if len(st.session_state.product_categories) >= 1:
|
||||
st.session_state.product_category = st.session_state.product_categories[0]
|
||||
|
||||
# Update list of projects.
|
||||
project_list_blob = bucket.blob("project_list.txt")
|
||||
project_list_blob.upload_from_string(
|
||||
json.dumps(st.session_state.product_categories)
|
||||
)
|
||||
st.rerun()
|
||||
|
||||
|
||||
def delete_file_from_gcs(file_name: str) -> None:
|
||||
"""Deletes a file from the GCS bucket.
|
||||
|
||||
This function deletes a file from the GCS bucket.
|
||||
It uses the 'storage_client' to get the 'blob' object for the file and
|
||||
then deletes the blob.
|
||||
|
||||
Args:
|
||||
file_name (str): The name of the file to delete.
|
||||
"""
|
||||
# Load and delete embeddings of deleted file.
|
||||
deleted_file_blob = bucket.blob(f"{st.session_state.product_category}/{file_name}")
|
||||
deleted_file_blob.delete()
|
||||
|
||||
# Load embeddings of the project
|
||||
project_embeddings = bucket.blob(
|
||||
st.session_state.product_category + "/embeddings.json"
|
||||
)
|
||||
stored_embedding_data = project_embeddings.download_as_string()
|
||||
embeddings_df = pd.DataFrame.from_dict(json.loads(stored_embedding_data))
|
||||
|
||||
# Remove deleted file from project embeddings.
|
||||
embeddings_df = embeddings_df.drop(
|
||||
embeddings_df[embeddings_df["file_name"] == file_name].index
|
||||
)
|
||||
embeddings_df.reset_index(inplace=True, drop=True)
|
||||
|
||||
# Update embeddings in GCS.
|
||||
bucket.blob(
|
||||
f"{st.session_state.product_category}/embeddings.json"
|
||||
).upload_from_string(embeddings_df.to_json(), "application/json")
|
||||
@@ -0,0 +1,15 @@
|
||||
backoff
|
||||
google-cloud-aiplatform
|
||||
numpy
|
||||
opencv-python
|
||||
opencv-python-headless
|
||||
pandas
|
||||
pillow
|
||||
PyPDF2
|
||||
dotenv
|
||||
python-docx
|
||||
pytomlpp
|
||||
streamlit
|
||||
streamlit-drawable-canvas
|
||||
vertexai
|
||||
google-cloud-storage==2.19.0
|
||||
+362
@@ -0,0 +1,362 @@
|
||||
"""
|
||||
This module provides functions for processing uploaded files, generating text
|
||||
embeddings,
|
||||
and managing data within the Google Cloud Storage (GCS) bucket.
|
||||
|
||||
This module:
|
||||
* Parses different file formats (CSV, text, Word, PDF), extracts
|
||||
text, splits it into chunks, and creates data packets.
|
||||
* Leverages `embedding_model_with_backoff` to embed text chunks.
|
||||
* Uploads processed data packets to a GCS bucket.
|
||||
* Stores embeddings alongside their associated metadata.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from PyPDF2 import PdfReader
|
||||
import aiohttp as cloud_function_call
|
||||
from app.pages_utils import insights
|
||||
from app.pages_utils.embedding_model import embedding_model_with_backoff
|
||||
from app.pages_utils.pages_config import GLOBAL_CFG
|
||||
import docx
|
||||
from dotenv import load_dotenv
|
||||
from google.cloud import storage
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import streamlit as st
|
||||
from streamlit.runtime.uploaded_file_manager import UploadedFile
|
||||
|
||||
load_dotenv()
|
||||
|
||||
PROJECT_ID = os.getenv("PROJECT_ID")
|
||||
LOCATION = os.getenv("LOCATION")
|
||||
|
||||
logging.basicConfig(format="%(levelname)s:%(message)s", level=logging.DEBUG)
|
||||
|
||||
# Define storage bucket
|
||||
storage_client = storage.Client(project=PROJECT_ID)
|
||||
bucket = storage_client.bucket(GLOBAL_CFG["bucket_name"])
|
||||
|
||||
|
||||
def get_chunks_iter(text: str, maxlength: int) -> list[str]:
|
||||
"""Gets the chunks of text from a string.
|
||||
|
||||
This function gets the chunks of text from a string.
|
||||
It splits the string into chunks of the specified maximum length and
|
||||
returns a list of the chunks.
|
||||
|
||||
Args:
|
||||
text (str): The string to get the chunks of.
|
||||
maxlength (int): The maximum length of the chunks.
|
||||
|
||||
Returns:
|
||||
list[str]: A list of the chunks of text.
|
||||
"""
|
||||
start = 0
|
||||
end = 0
|
||||
final_chunk = []
|
||||
while start + maxlength < len(text) and end != -1:
|
||||
end = text.rfind(" ", start, start + maxlength + 1)
|
||||
final_chunk.append(text[start:end])
|
||||
start = end + 1
|
||||
final_chunk.append(text[start:])
|
||||
return final_chunk
|
||||
|
||||
|
||||
def chunk_and_store_data(
|
||||
uploaded_file: UploadedFile,
|
||||
file_content: str,
|
||||
) -> list:
|
||||
"""Creates a data packet.
|
||||
|
||||
This function creates and returns a list of chunks from the file contents.
|
||||
|
||||
Args:
|
||||
uploaded_file: File like object from streamlit uploader.
|
||||
file_content (str): The contents of the file.
|
||||
|
||||
Returns:
|
||||
final_data (list[Any]): A list of data packets.
|
||||
"""
|
||||
# Creating a simple dictionary to store all information
|
||||
# (content and metadata) extracted from the document
|
||||
|
||||
# Return if empty or invalid file is found.
|
||||
if file_content == "":
|
||||
return []
|
||||
|
||||
final_data = []
|
||||
|
||||
# Split file into chunks and process each chunk in parallel.
|
||||
text_chunks = get_chunks_iter(file_content, 2000)
|
||||
for chunk_number, chunk_content in enumerate(text_chunks):
|
||||
data_packet = {}
|
||||
data_packet["file_name"] = uploaded_file.name
|
||||
data_packet["chunk_number"] = str(chunk_number)
|
||||
data_packet["content"] = chunk_content
|
||||
|
||||
# Append all chunks to final_data.
|
||||
final_data.append(data_packet)
|
||||
|
||||
return final_data
|
||||
|
||||
|
||||
async def add_embedding_col(pdf_data: pd.DataFrame) -> pd.DataFrame:
|
||||
"""Adds an 'embedding' column to the PDF data.
|
||||
|
||||
This function adds an 'embedding' column to the PDF data.
|
||||
It uses the 'apply' function to apply the 'generate_embeddings' function
|
||||
to each row of the 'content' column and returns the resulting DataFrame.
|
||||
|
||||
Args:
|
||||
pdf_data (pd.DataFrame): The PDF data.
|
||||
ti (int): The task index.
|
||||
|
||||
Returns:
|
||||
pd.DataFrame: The PDF data with the 'embedding' column.
|
||||
"""
|
||||
# Make request data payload
|
||||
pdf_content = json.dumps({"pdf_data": pdf_data["content"].to_json()})
|
||||
|
||||
async with cloud_function_call.ClientSession() as session:
|
||||
# URL for cloud function.
|
||||
url = f"""https://us-central1-{PROJECT_ID}.cloudfunctions.net/text-embedding"""
|
||||
|
||||
# Call cloud function to generate embeddings with data and headers.
|
||||
async with session.post(
|
||||
url,
|
||||
data=pdf_content,
|
||||
headers=st.session_state.headers,
|
||||
verify_ssl=False,
|
||||
) as embedding_response:
|
||||
# Process cloud function Response
|
||||
if embedding_response.status == 200:
|
||||
response = await embedding_response.text() # Read response
|
||||
# Extract text embeddings and convert to pd.Series.
|
||||
text_embeddings = pd.Series(json.loads(response)["embedding_column"][0])
|
||||
# Add embedding column
|
||||
pdf_data["embedding"] = text_embeddings
|
||||
|
||||
return pdf_data
|
||||
|
||||
|
||||
async def process_rows(df: pd.DataFrame, filename: str, header: list) -> pd.DataFrame:
|
||||
"""Processes the rows.
|
||||
|
||||
This function processes the rows.
|
||||
It iterates over the rows of the DataFrame and creates a data packet for
|
||||
each row.
|
||||
It then returns a DataFrame with the data packets.
|
||||
|
||||
Args:
|
||||
df (pd.DataFrame): The DataFrame to process.
|
||||
filename (str): The name of the file.
|
||||
header (list): The header of the file.
|
||||
|
||||
Returns:
|
||||
pd.DataFrame: A DataFrame with the data packets.
|
||||
"""
|
||||
final_data = []
|
||||
last_index = len(df)
|
||||
for i in range(last_index):
|
||||
chunk_content = ""
|
||||
for j, head in enumerate(header):
|
||||
chunk_content += f"{head} is {df.iloc[[i], [j]].squeeze()}. "
|
||||
data_packet = {}
|
||||
data_packet["file_name"] = filename
|
||||
data_packet["chunk_number"] = str(i + 1)
|
||||
data_packet["content"] = chunk_content
|
||||
|
||||
final_data.append(data_packet)
|
||||
pdf_data = pd.DataFrame.from_dict(final_data)
|
||||
return pdf_data
|
||||
|
||||
|
||||
async def csv_processing(
|
||||
df: pd.DataFrame,
|
||||
header: list,
|
||||
embeddings_df: pd.DataFrame,
|
||||
file: str,
|
||||
) -> None:
|
||||
"""Processes the CSV file.
|
||||
|
||||
This function processes the CSV file.
|
||||
It splits the DataFrame into chunks and processes each chunk in parallel.
|
||||
It then concatenates the results and uploads the resulting DataFrame to
|
||||
the GCS bucket.
|
||||
|
||||
Args:
|
||||
df (pd.DataFrame): The DataFrame to process.
|
||||
header (list): The header of the file.
|
||||
embeddings_df (pd.DataFrame): The DataFrame with the stored embeddings.
|
||||
file (str): The name of the file.
|
||||
"""
|
||||
pdf_data = pd.DataFrame()
|
||||
df_size = len(df)
|
||||
chunk_size = 100 # Default chunk size if df_size is below all thresholds
|
||||
chunk_sizes = {
|
||||
1_000_000: 100_000,
|
||||
100_000: 10_000,
|
||||
10_000: 1_000,
|
||||
}
|
||||
for threshold, size in chunk_sizes.items():
|
||||
if df_size > threshold:
|
||||
chunk_size = size
|
||||
|
||||
chunks = np.array_split(df, chunk_size)
|
||||
|
||||
# Parallel processing in stages
|
||||
processed_chunks = await asyncio.gather(
|
||||
*(process_rows(chunk, file, header) for _, chunk in enumerate(chunks))
|
||||
)
|
||||
typed_chunks = await asyncio.gather(
|
||||
*(
|
||||
asyncio.to_thread(
|
||||
chunk.assign(types=[type(x) for x in chunk["content"]]), chunk
|
||||
)
|
||||
for chunk in processed_chunks
|
||||
)
|
||||
)
|
||||
embedded_chunks = await asyncio.gather(
|
||||
*(add_embedding_col(chunk) for chunk in typed_chunks)
|
||||
)
|
||||
|
||||
# Combine, merge, deduplicate, and upload
|
||||
pdf_data = pd.concat(embedded_chunks + [embeddings_df])
|
||||
pdf_data = pdf_data.drop_duplicates(subset="content", keep="first").reset_index(
|
||||
drop=True
|
||||
)
|
||||
bucket.blob(
|
||||
f"{st.session_state.product_category}/embeddings.json"
|
||||
).upload_from_string(pdf_data.to_json(), "application/json")
|
||||
|
||||
|
||||
def load_file_content(
|
||||
uploaded_file: UploadedFile,
|
||||
uploaded_file_blob: storage.Blob,
|
||||
) -> Any:
|
||||
"""Loads and processes the content of various file types (text, docx, pdf).
|
||||
|
||||
Args:
|
||||
uploaded_file: The file to convert to data packets.
|
||||
uploaded_file_blob (optional): A Google Cloud Storage Blob object. If
|
||||
provided, the function will upload the file content to the blob.
|
||||
|
||||
Returns:
|
||||
The extracted text content of the file(s) as a single string.
|
||||
"""
|
||||
# Handle case if a text file has been uploaded.
|
||||
if uploaded_file.type == "text/plain":
|
||||
# Read and decode contents of the file.
|
||||
file_content = uploaded_file.read().decode("utf-8")
|
||||
file_content = file_content.replace("\n", " ")
|
||||
uploaded_file_blob.upload_from_string(
|
||||
file_content, content_type=uploaded_file.type
|
||||
)
|
||||
# Handle case when uploaded file is a document.
|
||||
elif uploaded_file.name.lower().endswith(".docx"):
|
||||
# Read and clean up contents of the document.
|
||||
doc = docx.Document(uploaded_file)
|
||||
file_content = ""
|
||||
for para in doc.paragraphs:
|
||||
file_content += para.text
|
||||
"\n".join(file_content)
|
||||
uploaded_file_blob.upload_from_string(
|
||||
file_content, content_type=uploaded_file.type
|
||||
)
|
||||
else:
|
||||
# Read and process contents of the pdf file.
|
||||
pdf_content = uploaded_file.read()
|
||||
uploaded_file_blob.upload_from_string(
|
||||
pdf_content, content_type=uploaded_file.type
|
||||
)
|
||||
# Extract pages from the pdf.
|
||||
reader = PdfReader(uploaded_file)
|
||||
num_pages = len(reader.pages)
|
||||
file_content = ""
|
||||
# Separately load text from each page of the pdf.
|
||||
for page_num in range(num_pages):
|
||||
page = reader.pages[page_num]
|
||||
pg = page.extract_text()
|
||||
file_content += pg
|
||||
return file_content
|
||||
|
||||
|
||||
def create_and_store_embeddings(uploaded_file: UploadedFile) -> None:
|
||||
"""Converts the file to data packets.
|
||||
|
||||
This function converts the file to data packets.
|
||||
It checks the file type and processes the file accordingly.
|
||||
It then uploads the resulting DataFrame to the GCS bucket.
|
||||
|
||||
Args:
|
||||
uploaded_file: The file to convert to data packets.
|
||||
"""
|
||||
with st.spinner("Uploading files..."):
|
||||
uploaded_file_blob = bucket.blob(
|
||||
f"{st.session_state.product_category}/{uploaded_file.name}"
|
||||
)
|
||||
|
||||
embeddings_df = insights.get_stored_embeddings_as_df()
|
||||
final_data = []
|
||||
|
||||
# Processing for csv/text files.
|
||||
if uploaded_file.type == "text/csv":
|
||||
# Read the csv file contents.
|
||||
df = pd.read_csv(uploaded_file)
|
||||
uploaded_file_blob.upload_from_string(df.to_csv(), "text/csv")
|
||||
|
||||
# Return if file is empty or contents cannot be read.
|
||||
if df.empty:
|
||||
return
|
||||
|
||||
# Create a list of csv file columns.
|
||||
header = []
|
||||
for col in df.columns:
|
||||
header.append(col)
|
||||
|
||||
# Create embeddings and store contents of the csv file
|
||||
# to the GCS bucket.
|
||||
with st.spinner("Processing csv...this might take some time..."):
|
||||
asyncio.run(
|
||||
csv_processing(df, header, embeddings_df, uploaded_file.name)
|
||||
)
|
||||
return
|
||||
|
||||
file_content = load_file_content(uploaded_file, uploaded_file_blob)
|
||||
# Append processed content from the page to final data.
|
||||
final_data = chunk_and_store_data(
|
||||
uploaded_file=uploaded_file,
|
||||
file_content=file_content,
|
||||
)
|
||||
if len(final_data) == 0:
|
||||
return
|
||||
# Stores the embeddings in the GCS bucket.
|
||||
with st.spinner("Storing Embeddings"):
|
||||
# Create a dataframe from final chunked data.
|
||||
pdf_data = pd.DataFrame.from_dict(final_data)
|
||||
pdf_data.reset_index(inplace=True, drop=True)
|
||||
|
||||
# Add datatype column to df.
|
||||
pdf_data["types"] = [type(x) for x in pdf_data["content"]]
|
||||
|
||||
# Add embedding column to df for text embeddings.
|
||||
pdf_data["embedding"] = pdf_data["content"].apply(
|
||||
lambda x: embedding_model_with_backoff([x])
|
||||
)
|
||||
pdf_data["embedding"] = pdf_data.embedding.apply(np.array)
|
||||
|
||||
# Concatenate the data of newly uploaded files with that of
|
||||
# existing file embeddings
|
||||
pdf_data = pd.concat([embeddings_df, pdf_data])
|
||||
pdf_data = pdf_data.drop_duplicates(subset=["content"], keep="first")
|
||||
pdf_data.reset_index(inplace=True, drop=True)
|
||||
|
||||
# Upload newly created embeddings to gcs
|
||||
bucket.blob(
|
||||
f"{st.session_state.product_category}/embeddings.json"
|
||||
).upload_from_string(pdf_data.to_json(), "application/json")
|
||||
@@ -0,0 +1,176 @@
|
||||
"""
|
||||
Common utilities for the project. This includes:
|
||||
* session state initialization.
|
||||
* project selection.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from app.pages_utils.pages_config import GLOBAL_CFG
|
||||
from google.cloud import storage
|
||||
import streamlit as st
|
||||
from vertexai import generative_models
|
||||
|
||||
PROJECT_ID = os.getenv("PROJECT_ID")
|
||||
LOCATION = os.getenv("LOCATION")
|
||||
|
||||
# Define storage bucket
|
||||
storage_client = storage.Client(project=PROJECT_ID)
|
||||
bucket = storage_client.bucket(GLOBAL_CFG["bucket_name"])
|
||||
|
||||
|
||||
def display_projects() -> None:
|
||||
"""Displays the list of projects and allows the user to select one.
|
||||
|
||||
Args:
|
||||
None
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
st.session_state.product_category = st.selectbox(
|
||||
"Select a project", st.session_state.product_categories
|
||||
)
|
||||
st.session_state.product_categories.remove(st.session_state.product_category)
|
||||
st.session_state.product_categories.insert(0, st.session_state.product_category)
|
||||
if st.session_state.previous_product_category != st.session_state.product_category:
|
||||
initialize_all_session_state(reinitialize=True)
|
||||
st.session_state.previous_product_category = st.session_state.product_category
|
||||
st.rerun()
|
||||
|
||||
|
||||
def initialize_all_session_state(reinitialize: bool = False) -> None:
|
||||
"""Initializes all the session states used in the app.
|
||||
|
||||
Args:
|
||||
reinitialize (optional, bool):
|
||||
Indicated if the session state is being reinitialized
|
||||
or being initialized for the first time.
|
||||
(This value is important to indicate that the value of
|
||||
the selected project has been updated. If it is set to false
|
||||
then no modification is made to the session state).
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
# Get lists of projects in the application.
|
||||
project_list_blob = bucket.blob("project_list.txt")
|
||||
project_list = json.loads(project_list_blob.download_as_string())
|
||||
|
||||
# Initialize default values for the session state.
|
||||
session_state_defaults: dict[str, Any] = {
|
||||
"product_categories": project_list,
|
||||
"new_product_category_added": None,
|
||||
"previous_product_category": None,
|
||||
"text_edit_prompt": None,
|
||||
"headers": {"Content-Type": "application/json"},
|
||||
"update_text_btn": None,
|
||||
"uploaded_files": None,
|
||||
"rag_search_term": None,
|
||||
"rag_answers_gen": False,
|
||||
"rag_answer": None,
|
||||
"rag_answer_references": None,
|
||||
"insights_suggestion": None,
|
||||
"insights_placeholder": "",
|
||||
"suggestion_first_time": 1,
|
||||
"processed_data_list": [],
|
||||
"query_vectors": [],
|
||||
"embeddings_df": None,
|
||||
"temp_suggestions": None,
|
||||
"assorted_prod_title": None,
|
||||
"assorted_prod_content": [],
|
||||
"email_gen": False,
|
||||
"create_product": False,
|
||||
"modifying": False,
|
||||
"custom_prompt": "",
|
||||
"feature_suggestions": None,
|
||||
"selected_titles": [],
|
||||
"saved_titles": [],
|
||||
"selected_prompt": None,
|
||||
"product_gen_image": None,
|
||||
"features_generated": False,
|
||||
"generated_points": None,
|
||||
"content_generated": False,
|
||||
"product_content": None,
|
||||
"image_to_edit": -1,
|
||||
"generate_images": False,
|
||||
"image_prompt": None,
|
||||
"image_file_prefix": "uploaded_image",
|
||||
"email_image": None,
|
||||
"email_prompt": "High SPF",
|
||||
"num_drafts": None,
|
||||
"email_text": None,
|
||||
"generated_image": None,
|
||||
"mask_image": None,
|
||||
"edit_suggestion": False,
|
||||
"suggested_images": None,
|
||||
"uploaded_img": False,
|
||||
"start_editing": False,
|
||||
"text_to_edit": None,
|
||||
"content_edited": None,
|
||||
"row": None,
|
||||
"edited_content": None,
|
||||
"col": None,
|
||||
"generated_response": None,
|
||||
"draft_elements": None,
|
||||
"chosen_titles": [],
|
||||
"buffer": None,
|
||||
"save_edited_image": None,
|
||||
"email_files": [],
|
||||
"image_edit_col": None,
|
||||
"image_edit_row": None,
|
||||
"bg_editing": False,
|
||||
}
|
||||
|
||||
for key, value in session_state_defaults.items():
|
||||
if (
|
||||
reinitialize is False and key not in st.session_state
|
||||
) or reinitialize is True:
|
||||
st.session_state[key] = value
|
||||
|
||||
if "product_category" not in st.session_state:
|
||||
st.session_state.product_category = st.session_state.product_categories[0]
|
||||
|
||||
st.session_state.generation_config = generative_models.GenerationConfig(
|
||||
max_output_tokens=8192,
|
||||
temperature=0.001,
|
||||
top_p=1,
|
||||
)
|
||||
|
||||
|
||||
def page_setup(page_cfg: dict) -> None:
|
||||
"""
|
||||
This function initializes the page configuration and applies custom styles.
|
||||
|
||||
Args:
|
||||
page_cfg (dict): A dictionary containing the configuration for the
|
||||
page.
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
|
||||
# Set the page configuration
|
||||
st.set_page_config(
|
||||
page_title=page_cfg["page_title"], page_icon=page_cfg["page_icon"]
|
||||
)
|
||||
|
||||
# Initialize session state for the project if it does not exist.
|
||||
if (
|
||||
"initialize_session_state" not in st.session_state
|
||||
or st.session_state.initialize_session_state is False
|
||||
):
|
||||
initialize_all_session_state()
|
||||
st.session_state.initialize_session_state = True
|
||||
# Apply the sidebar style
|
||||
load_css("app/css/sidebar_styles.css")
|
||||
|
||||
|
||||
def load_css(css_file_path: str) -> None:
|
||||
"""
|
||||
Load css from the given filepath.
|
||||
"""
|
||||
with open(css_file_path, encoding="utf-8") as f:
|
||||
st.markdown(f"<style>{f.read()}</style>", unsafe_allow_html=True)
|
||||
@@ -0,0 +1 @@
|
||||
streamlit
|
||||
+63
@@ -0,0 +1,63 @@
|
||||
"""
|
||||
Cloud Function for getting text response from Gemini API.
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from dotenv import load_dotenv
|
||||
import functions_framework
|
||||
from vertexai.preview import generative_models
|
||||
from vertexai.preview.generative_models import GenerativeModel
|
||||
|
||||
load_dotenv()
|
||||
|
||||
PROJECT_ID = os.getenv("PROJECT_ID")
|
||||
LOCATION = os.getenv("LOCATION")
|
||||
|
||||
|
||||
def generate_text(prompt: str) -> str:
|
||||
"""Generates text using the Gemini-Pro model.
|
||||
|
||||
Args:
|
||||
prompt: The text prompt to use for generation.
|
||||
|
||||
Returns:
|
||||
The generated text.
|
||||
"""
|
||||
model = GenerativeModel("gemini-2.0-flash")
|
||||
|
||||
generation_config = generative_models.GenerationConfig(
|
||||
max_output_tokens=8192,
|
||||
temperature=0.001,
|
||||
top_p=1,
|
||||
)
|
||||
|
||||
response = model.generate_content(
|
||||
prompt,
|
||||
generation_config=generation_config,
|
||||
)
|
||||
return response.text
|
||||
|
||||
|
||||
@functions_framework.http
|
||||
def get_llm_response(request: Any) -> dict | tuple:
|
||||
"""HTTP Cloud Function that generates text using the Gemini-Pro model.
|
||||
|
||||
Args:
|
||||
request (flask.Request): The request object.
|
||||
<http://flask.palletsprojects.com/en/1.1.x/api/#incoming-request-data>
|
||||
|
||||
Returns:
|
||||
The response text, or any set of values that can be turned into a
|
||||
Response object using `make_response`
|
||||
<http://flask.palletsprojects.com/en/1.1.x/api/#flask.make_response>.
|
||||
"""
|
||||
request_json: dict = request.get_json(silent=True)
|
||||
if not request_json or "text_prompt" not in request_json:
|
||||
return {"error": "Request body must contain 'text_prompt' field."}, 400
|
||||
|
||||
text_prompt = request_json["text_prompt"]
|
||||
generated_text = generate_text(text_prompt)
|
||||
|
||||
return {"generated_text": generated_text}
|
||||
+3
@@ -0,0 +1,3 @@
|
||||
functions-framework==3.*
|
||||
google-cloud-aiplatform
|
||||
dotenv
|
||||
+50
@@ -0,0 +1,50 @@
|
||||
"""
|
||||
Cloud function to make calls to Imagen API.
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from dotenv import load_dotenv
|
||||
import functions_framework
|
||||
import vertexai
|
||||
from vertexai.preview.vision_models import ImageGenerationModel
|
||||
|
||||
load_dotenv()
|
||||
|
||||
PROJECT_ID = os.getenv("PROJECT_ID")
|
||||
LOCATION = os.getenv("LOCATION")
|
||||
|
||||
|
||||
def image_generation(prompt: str) -> bytes:
|
||||
"""Generates images based on given prompt using image generation model.
|
||||
|
||||
Args:
|
||||
prompt (str): Prompt for generating image.
|
||||
|
||||
Returns:
|
||||
bytes: The generated image as raw bytes.
|
||||
"""
|
||||
vertexai.init(project=PROJECT_ID, location=LOCATION)
|
||||
model = ImageGenerationModel.from_pretrained("imagegeneration@006")
|
||||
image = model.generate_images(
|
||||
prompt=prompt,
|
||||
number_of_images=1,
|
||||
language="en",
|
||||
aspect_ratio="1:1",
|
||||
)[0]
|
||||
return image._loaded_bytes # pylint: disable=protected-access
|
||||
|
||||
|
||||
@functions_framework.http
|
||||
def get_images(request: Any) -> bytes:
|
||||
"""Invokes image generation call.
|
||||
|
||||
Args:
|
||||
request: Data for image generation from the calling function.
|
||||
|
||||
Returns:
|
||||
Response: A Flask Response object containing the generated image.
|
||||
"""
|
||||
request_json: dict = request.get_json(silent=True)
|
||||
return image_generation(request_json["img_prompt"])
|
||||
+4
@@ -0,0 +1,4 @@
|
||||
functions-framework==3.*
|
||||
dotenv
|
||||
vertexai
|
||||
|
||||
+86
@@ -0,0 +1,86 @@
|
||||
"""
|
||||
Cloud function to generate embedding of given file.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from dotenv import load_dotenv
|
||||
import functions_framework
|
||||
from vertexai.preview.language_models import TextEmbeddingModel
|
||||
|
||||
load_dotenv()
|
||||
|
||||
|
||||
PROJECT_ID = os.getenv("PROJECT_ID")
|
||||
LOCATION = os.getenv("LOCATION")
|
||||
|
||||
|
||||
embedding_model = TextEmbeddingModel.from_pretrained("text-embedding-005")
|
||||
|
||||
|
||||
def get_embeddings(instances: list[str]) -> list[list[float]]:
|
||||
"""
|
||||
Generates embeddings for given text.
|
||||
|
||||
Args:
|
||||
instance (list[str]):
|
||||
Text to convert to embeddings.
|
||||
|
||||
Returns:
|
||||
embeddings (list):
|
||||
values of embeddings.
|
||||
"""
|
||||
|
||||
embeddings = embedding_model.get_embeddings(instances)
|
||||
return [embedding.values for embedding in embeddings]
|
||||
|
||||
|
||||
def generate_embeddings(pdf_data: dict) -> dict:
|
||||
"""
|
||||
Extracts content from pdf_data for creating embeddings.
|
||||
|
||||
Args:
|
||||
pdf_data (dict): file data to be processed.
|
||||
"""
|
||||
instances = []
|
||||
values = []
|
||||
|
||||
batch_size = 10
|
||||
iterate = 0
|
||||
|
||||
for content in pdf_data.values():
|
||||
instances.append(content)
|
||||
iterate += 1
|
||||
|
||||
if iterate % batch_size == 0 or iterate == len(pdf_data):
|
||||
embeddings = get_embeddings(instances)
|
||||
values.append(embeddings)
|
||||
|
||||
instances = []
|
||||
|
||||
response_json = json.dumps({"embedding_column": values})
|
||||
response = json.loads(response_json)
|
||||
return response
|
||||
|
||||
|
||||
@functions_framework.http
|
||||
def get_text_embeddings(request: Any) -> tuple[dict[str, str], int]:
|
||||
"""
|
||||
Processes request for generating embeddings.
|
||||
|
||||
Args:
|
||||
request:
|
||||
Data for conversion to embeddings with
|
||||
headers by the calling func.
|
||||
Returns:
|
||||
embeddings (dict):
|
||||
generated embeddings
|
||||
"""
|
||||
request_json = request.get_json(silent=True)
|
||||
if not request_json or "pdf_data" not in request_json:
|
||||
return {"error": "Request body must contain 'pdf_data' field."}, 400
|
||||
pdf_data = request_json["pdf_data"]
|
||||
embeddings = generate_embeddings(pdf_data)
|
||||
return embeddings, 200
|
||||
+3
@@ -0,0 +1,3 @@
|
||||
functions-framework==3.*
|
||||
dotenv
|
||||
vertexai
|
||||
@@ -0,0 +1,102 @@
|
||||
#!/bin/bash
|
||||
source .env
|
||||
|
||||
echo "se: $YOUR_EMAIL"
|
||||
echo "pn : $PROJECT_NUMBER"
|
||||
echo "pID : $PROJECT_ID"
|
||||
|
||||
gcloud init --account "$YOUR_EMAIL" --project "$PROJECT_ID"
|
||||
gcloud auth application-default set-quota-project "$PROJECT_ID"
|
||||
gcloud config set project "$PROJECT_ID"
|
||||
|
||||
SERVICE_ACCOUNT="retail-accelerating-prod-i-982@$PROJECT_ID.iam.gserviceaccount.com"
|
||||
gcloud iam service-accounts add-iam-policy-binding "$SERVICE_ACCOUNT" --member "user:$YOUR_EMAIL" --role roles/iam.serviceAccountUser
|
||||
|
||||
gcloud functions deploy imagen-call \
|
||||
--allow-unauthenticated \
|
||||
--service-account="$SERVICE_ACCOUNT" \
|
||||
--run-service-account="$SERVICE_ACCOUNT" \
|
||||
--gen2 \
|
||||
--runtime=python311 \
|
||||
--region="$REGION" \
|
||||
--source=./cloud_functions/imagen_call \
|
||||
--entry-point=get_images \
|
||||
--trigger-http \
|
||||
--set-env-vars location="$LOCATION" \
|
||||
--set-env-vars project_id="$PROJECT_ID" \
|
||||
--set-env-vars MEMORY=512MB >cloud_fn_1
|
||||
|
||||
file="cloud_fn_1"
|
||||
previous_data=""
|
||||
while IFS= read -r line; do
|
||||
for word in $line; do
|
||||
if [ "$previous_data" == "url:" ]; then
|
||||
imagen_call_url="$word"
|
||||
fi
|
||||
previous_data="$word"
|
||||
done
|
||||
done <"$file"
|
||||
echo "Imagen Call URL: $imagen_call_url" >cloud_functions_urls
|
||||
|
||||
gcloud functions deploy gemini-call \
|
||||
--allow-unauthenticated \
|
||||
--service-account="$SERVICE_ACCOUNT" \
|
||||
--run-service-account="$SERVICE_ACCOUNT" \
|
||||
--gen2 \
|
||||
--runtime=python311 \
|
||||
--region="$REGION" \
|
||||
--source=./cloud_functions/gemini-call \
|
||||
--entry-point=get_llm_response \
|
||||
--trigger-http \
|
||||
--set-env-vars location="$LOCATION" \
|
||||
--set-env-vars project_id="$PROJECT_ID" \
|
||||
--set-env-vars MEMORY=512MB >cloud_fn_1
|
||||
|
||||
while IFS= read -r line; do
|
||||
for word in $line; do
|
||||
if [ "$previous_data" == "url:" ]; then
|
||||
text_bison_url="$word"
|
||||
fi
|
||||
previous_data="$word"
|
||||
done
|
||||
done <"$file"
|
||||
echo "Text Bison Call URL: $text_bison_url" >>cloud_functions_urls
|
||||
|
||||
gcloud functions deploy text-embedding \
|
||||
--allow-unauthenticated \
|
||||
--service-account="$SERVICE_ACCOUNT" \
|
||||
--run-service-account="$SERVICE_ACCOUNT" \
|
||||
--gen2 \
|
||||
--runtime=python311 \
|
||||
--region="$REGION" \
|
||||
--source=./cloud_functions/text-embedding \
|
||||
--entry-point=get_text_embeddings \
|
||||
--trigger-http \
|
||||
--set-env-vars location="$LOCATION" \
|
||||
--set-env-vars project_id="$PROJECT_ID" \
|
||||
--set-env-vars MEMORY=512MB >cloud_fn_1
|
||||
|
||||
while IFS= read -r line; do
|
||||
for word in $line; do
|
||||
if [ "$previous_data" == "url:" ]; then
|
||||
text_embedding_url="$word"
|
||||
fi
|
||||
previous_data="$word"
|
||||
done
|
||||
done <"$file"
|
||||
echo "Text Embedding URL: $text_embedding_url" >>cloud_functions_urls
|
||||
rm cloud_fn_1
|
||||
|
||||
# Set project ID, region, and service name (modify as needed)
|
||||
SERVICE_NAME="accelerating-product-innovation"
|
||||
|
||||
# Build the container image
|
||||
gcloud builds submit --tag "gcr.io/$PROJECT_ID/$SERVICE_NAME" .
|
||||
|
||||
# Deploy the image to Cloud Run
|
||||
gcloud run deploy "$SERVICE_NAME" \
|
||||
--image "gcr.io/$PROJECT_ID/$SERVICE_NAME" \
|
||||
--platform managed \
|
||||
--port 8080 \
|
||||
--region "$REGION" \
|
||||
--allow-unauthenticated
|
||||
@@ -0,0 +1,7 @@
|
||||
# 🚀 End-to-End Gen AI App Starter Pack 🚀
|
||||
|
||||
✨ **The e2e-gen-ai-app-starter-pack has moved!** ✨
|
||||
|
||||
Find the new and improved version at: [https://github.com/GoogleCloudPlatform/agent-starter-pack](https://github.com/GoogleCloudPlatform/agent-starter-pack)
|
||||
|
||||
This project represents the next evolution of the `e2e-gen-ai-app-starter-pack`.
|
||||
@@ -0,0 +1,10 @@
|
||||
FROM python:3.13
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
COPY requirements.txt .
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
COPY . .
|
||||
|
||||
CMD ["streamlit", "run", "home.py", "--server.enableCORS", "false", "--browser.serverAddress", "0.0.0.0", "--browser.gatherUsageStats", "false", "--server.port", "8080"]
|
||||
@@ -0,0 +1,128 @@
|
||||
# Finvest Spanner Demo App
|
||||
|
||||
**Authors:** [Anirban Bagchi](https://github.com/anirbanbagchi1979) and [Derek Downey](https://github.com/dtest)
|
||||
|
||||
<img align="right" style="padding-left: 10px;" src="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/finance-advisor-spanner/images/Finvest-white.jpg" width="35%" alt="Finvest Logo">
|
||||
|
||||
Consider a modern financial services company where I am a financial advisor. Finding the right financial investments can be challenging because of the complex nature of investments from structured data such as expense ratios, fund returns, to complex data such as asset holdings, their industry sectors, and more unstructured data, such as investment philosophy and client's investment goals. Let me show you how Spanner makes this process easy by combining these diverse data structures into a single multi-model platform.
|
||||
|
||||
The client wants me to find assets for funds in North America and Europe that invest in derivatives. I select North America and Europe and put in derivatives as my search term. Spanner runs a relational and text search to return a list of funds.
|
||||
|
||||
Next, the client wants to narrow this list to specific fund managers. I don't know the exact name, so I put in Liz Peters, and Spanner performs a fuzzy match(Full Text Search - Substring Match) of the name Liz Peters to find funds managed by Elizabeth Peterson.
|
||||
|
||||
Among these funds, the client wants to choose from socially responsible funds. Next, I check the box for vector search, and now I can see ESG funds because Spanner performed a KNN vector search to match the search term "socially responsible" with "environmental, social and governance".
|
||||
|
||||
Finally, before I recommend a fund, I also want to check the exposure to a particular sector. This can be complex because funds can invest in other funds, called fund of funds which makes it hard to compute this. Spanner performs a graph search using this asset knowledge graph. By traversing the funds and its holdings which could also be funds and their holdings, Spanner can compute the client's exposure to a particular sector. I can see the funds that have exposure of 20% or more in the technology sector.
|
||||
|
||||
With the power of Spanner's multimodel support, I can run complex workloads on a single database for relational, analytical, text and vector use cases with virtually unlimited scale, five nines of availability—including enterprise security and governance for mission critical workloads.
|
||||
|
||||
This demo highlights [Spanner](https://cloud.google.com/spanner), integration with [Vertex AI LLMs](https://cloud.google.com/model-garden?hl=en) for both embeddings and text completion models. You will learn how Spanner can help with use cases where you run Full Text Search, Approximate Nearest Neighbor search and vector similarity search.
|
||||
|
||||
## Tech Stack
|
||||
|
||||
The Finvest Spanner demo application was built using:
|
||||
|
||||
- [Spanner](https://cloud.google.com/spanner)
|
||||
- [Vertex AI](https://cloud.google.com/vertex-ai?hl=en) LLMs ([text-embedding-005](https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/text-embeddings) )
|
||||
- [Cloud Run](https://cloud.google.com/run)
|
||||
- [Dataflow](https://cloud.google.com/dataflow?)
|
||||
- [Streamlit](https://streamlit.io/)
|
||||
|
||||
## Deploying the Finvest Spanner Demo Application
|
||||
|
||||
1. Login to the [Google Cloud Console](https://console.cloud.google.com/).
|
||||
|
||||
2. [Create a new project](https://developers.google.com/maps/documentation/places/web-service/cloud-setup) to host the demo and isolate it from other resources in your account.
|
||||
|
||||
3. [Switch](https://cloud.google.com/resource-manager/docs/creating-managing-projects#identifying_projects) to your new project.
|
||||
|
||||
4. [Activate Cloud Shell](https://cloud.google.com/shell/docs/using-cloud-shell) and confirm your project by running the following commands. Click **Authorize** if prompted.
|
||||
|
||||
```bash
|
||||
gcloud auth list
|
||||
gcloud config list project
|
||||
```
|
||||
|
||||
5. Clone this repository and navigate to the project root:
|
||||
|
||||
```bash
|
||||
cd
|
||||
git clone https://github.com/GoogleCloudPlatform/generative-ai.git
|
||||
cd generative-ai/gemini/sample-apps/finance-advisor-spanner/
|
||||
```
|
||||
|
||||
6. Create a Spanner instance
|
||||
<https://console.cloud.google.com/spanner/instances/new>
|
||||
|
||||
> Note the instance Name
|
||||
|
||||
7. Import the data into the Spanner instance
|
||||
<https://cloud.google.com/spanner/docs/import#import-database>
|
||||
The bucket which has the Spanner export is in this public GCS Bucket
|
||||
|
||||
`https://storage.googleapis.com/github-repo/generative-ai/sample-apps/finance-advisor-spanner/spanner-fts-mf-data-export/`
|
||||
|
||||
> Note the Database Name
|
||||
|
||||
The import process will run and import the database into a new Spanner database.
|
||||
|
||||
8. Run Additional DDL statements for the database to have all the necessary components.
|
||||
The DDL statements are in [Schema-Operations.sql](./Schema-Operations.sql) file in this directory.
|
||||
|
||||
Change the endpoint as per your project and the spanner instance location
|
||||
|
||||
```sql
|
||||
ALTER MODEL EmbeddingsModel SET OPTIONS (
|
||||
endpoint = '//aiplatform.googleapis.com/projects/'YOUR PROJECT ID HERE'/locations/'YOUR SPANNER INSTANCE LOCATION HERE'/publishers/google/models/text-embedding-005'
|
||||
)
|
||||
;
|
||||
```
|
||||
|
||||
Next run the rest of DDL statements without any change
|
||||
|
||||
9. In Cloud Shell:
|
||||
|
||||
Open `.env` file in the same directory using vi or other Editor
|
||||
|
||||
Edit the following fields with the instance name from Step 6 and database name from Step 7
|
||||
|
||||
```bash
|
||||
instance_id='YOUR INSTANCE ID'
|
||||
database_id='YOUR DATABASE ID'
|
||||
```
|
||||
|
||||
10. Now Build & Deploy the application:
|
||||
|
||||
Build:
|
||||
|
||||
```bash
|
||||
gcloud builds submit --tag gcr.io/'YOUR PROJECT ID HERE'/finance-advisor-app
|
||||
```
|
||||
|
||||
Deploy:
|
||||
|
||||
```bash
|
||||
gcloud run deploy finance-advisor-app --image gcr.io/'YOUR PROJECT ID HERE'/finance-advisor-app --platform managed --region 'YOUR SPANNER REGION' --allow-unauthenticated
|
||||
```
|
||||
|
||||
### Troubleshooting
|
||||
|
||||
### Frontend
|
||||
|
||||
The frontend application is Streamlit running on CloudRun
|
||||
|
||||
## Purpose and Extensibility
|
||||
|
||||
The purpose of this repository is to help you provision an isolated demo environment that highlights the Full Text Search, Semantic Search and Graph capabilities of Spanner. While the ideas in this repository can be extended for many real-world use cases, the demo code itself is overly permissive and has not been hardened for security or reliability. The sample code in this repository is provided on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, and it should NOT be used for production use cases without doing your own testing and security hardening.
|
||||
|
||||
## Clean Up
|
||||
|
||||
Be sure to delete the resources you no longer need when you're done with the demo. If you created a new project for the lab as recommended, you can delete the whole project using the command below in your Cloud Shell session (NOT the pgadmin VM).
|
||||
|
||||
**DANGER: Be sure to set PROJECT_ID to the correct project, and run this command ONLY if you are SURE there is nothing in the project that you might still need. This command will permanently destroy everything in the project.**
|
||||
|
||||
```bash
|
||||
# Set your project id
|
||||
PROJECT_ID='YOUR PROJECT ID HERE'
|
||||
gcloud projects delete ${PROJECT_ID}
|
||||
```
|
||||
@@ -0,0 +1,136 @@
|
||||
|
||||
ALTER MODEL EmbeddingsModel SET OPTIONS (
|
||||
endpoint = '//aiplatform.googleapis.com/projects/<project-name>/locations/<location>/publishers/google/models/text-embedding-005'
|
||||
)
|
||||
;
|
||||
ALTER TABLE EU_MutualFunds ADD COLUMN fund_name_Tokens TOKENLIST AS (TOKENIZE_FULLTEXT(fund_name)) HIDDEN;
|
||||
ALTER TABLE EU_MutualFunds ADD COLUMN category_Tokens TOKENLIST AS (TOKENIZE_FULLTEXT(category)) HIDDEN;
|
||||
ALTER TABLE EU_MutualFunds ADD COLUMN investment_strategy_Tokens TOKENLIST AS (TOKENIZE_FULLTEXT(investment_strategy)) HIDDEN;
|
||||
ALTER TABLE EU_MutualFunds ADD COLUMN investment_managers_Tokens TOKENLIST AS (TOKENIZE_FULLTEXT(investment_managers)) HIDDEN;
|
||||
ALTER TABLE EU_MutualFunds ADD COLUMN fund_benchmark_Tokens TOKENLIST AS (TOKENIZE_FULLTEXT(fund_benchmark)) HIDDEN;
|
||||
ALTER TABLE EU_MutualFunds ADD COLUMN morningstar_benchmark_Tokens TOKENLIST AS (TOKENIZE_FULLTEXT(morningstar_benchmark)) HIDDEN;
|
||||
ALTER TABLE EU_MutualFunds ADD COLUMN top5_regions_Tokens TOKENLIST AS (TOKENIZE_FULLTEXT(top5_regions)) HIDDEN;
|
||||
ALTER TABLE EU_MutualFunds ADD COLUMN top5_holdings_Tokens TOKENLIST AS (TOKENIZE_FULLTEXT(top5_holdings)) HIDDEN;
|
||||
ALTER TABLE EU_MutualFunds ADD COLUMN investment_managers_Substring_Tokens TOKENLIST AS (TOKENIZE_SUBSTRING(investment_managers)) HIDDEN;
|
||||
ALTER TABLE
|
||||
EU_MutualFunds ADD COLUMN investment_managers_Substring_Tokens_NGRAM TOKENLIST AS ( TOKENIZE_SUBSTRING(investment_managers,
|
||||
ngram_size_min=>2,
|
||||
ngram_size_max=>3,
|
||||
relative_search_types=>["word_prefix",
|
||||
"word_suffix"])) HIDDEN;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1958 ;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2008;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2004;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1989;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2002;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2014;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2019;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2005;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1988;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2006;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1987;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2007;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1992;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1974;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2011;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1996;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2018;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1941;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1972;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1993;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2013;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1991;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2010;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1997;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2001;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2015;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1934;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1985;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1990;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2017;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1998;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1999;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2012;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1984;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1995;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2009;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2003;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1994;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1973;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1981;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2016;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2020;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 2000;
|
||||
UPDATE EU_MutualFunds SET investment_strategy_Embedding_vector = investment_strategy_Embedding WHERE investment_strategy_Embedding is not NULL and EXTRACT(YEAR from inception_date) = 1992;
|
||||
CREATE SEARCH INDEX
|
||||
category_Tokens_IDX
|
||||
ON
|
||||
EU_MutualFunds(category_Tokens);
|
||||
CREATE SEARCH INDEX
|
||||
fund_benchmark_Tokens_IDX
|
||||
ON
|
||||
EU_MutualFunds(fund_benchmark_Tokens);
|
||||
CREATE SEARCH INDEX
|
||||
fund_name_Tokens_IDX
|
||||
ON
|
||||
EU_MutualFunds(fund_name_Tokens);
|
||||
CREATE SEARCH INDEX
|
||||
investment_managers_Tokens_IDX
|
||||
ON
|
||||
EU_MutualFunds(investment_managers_Tokens);
|
||||
CREATE SEARCH INDEX
|
||||
investment_strategy_Tokens_IDX
|
||||
ON
|
||||
EU_MutualFunds(investment_strategy_Tokens);
|
||||
CREATE SEARCH INDEX
|
||||
morningstar_benchmark_Tokens_IDX
|
||||
ON
|
||||
EU_MutualFunds(morningstar_benchmark_Tokens);
|
||||
CREATE SEARCH INDEX
|
||||
top5_holdings_Tokens_IDX
|
||||
ON
|
||||
EU_MutualFunds(top5_holdings_Tokens);
|
||||
CREATE SEARCH INDEX
|
||||
top5_regions_Tokens_IDX
|
||||
ON
|
||||
EU_MutualFunds(top5_regions_Tokens);
|
||||
CREATE SEARCH INDEX
|
||||
investment_managers_Substring_Tokens_IDX
|
||||
ON
|
||||
EU_MutualFunds(investment_managers_Substring_Tokens);
|
||||
CREATE SEARCH INDEX
|
||||
investment_managers_Substring_investment_Strategy_Tokens_Combo_IDX
|
||||
ON
|
||||
EU_MutualFunds(investment_managers_Substring_Tokens,
|
||||
investment_strategy_Tokens);
|
||||
CREATE SEARCH INDEX
|
||||
investment_managers_Substring_NgRAM_investment_Strategy_Tokens_Combo_IDX
|
||||
ON
|
||||
EU_MutualFunds(investment_strategy_Tokens,
|
||||
investment_managers_Substring_Tokens_NGRAM);
|
||||
CREATE VECTOR INDEX
|
||||
InvestmentStrategyEmbeddingIndex
|
||||
ON
|
||||
EU_MutualFunds(investment_strategy_Embedding_vector)
|
||||
WHERE
|
||||
investment_strategy_Embedding_vector IS NOT NULL OPTIONS ( tree_depth = 2,
|
||||
num_leaves = 40,
|
||||
distance_type = 'EUCLIDEAN' );
|
||||
CREATE SEARCH INDEX
|
||||
investment_managers_Substring_Tokens_with_vectors_NGRAM_IDX
|
||||
ON
|
||||
EU_MutualFunds(investment_managers_Substring_Tokens_NGRAM) STORING (investment_strategy_Embedding_vector);
|
||||
CREATE OR REPLACE PROPERTY GRAPH FundGraph NODE TABLES( Companies AS Company DEFAULT LABEL PROPERTIES ALL COLUMNS,
|
||||
EU_MutualFunds AS Fund DEFAULT LABEL PROPERTIES ALL COLUMNS EXCEPT (_Injected_SearchUid,
|
||||
_Injected_VectorIndex_InvestmentStrategyEmbeddingIndex_FP8,
|
||||
_Injected_VectorIndex_InvestmentStrategyEmbeddingIndex_LeafId),
|
||||
Sectors AS Sector DEFAULT LABEL PROPERTIES ALL COLUMNS ) EDGE TABLES( FundHoldsCompany SOURCE KEY(NewMFSequence)
|
||||
REFERENCES
|
||||
Fund(NewMFSequence) DESTINATION KEY(CompanySeq)
|
||||
REFERENCES
|
||||
Company(CompanySeq) LABEL Holds PROPERTIES ALL COLUMNS,
|
||||
CompanyBelongsSector SOURCE KEY(CompanySeq)
|
||||
REFERENCES
|
||||
Company(CompanySeq) DESTINATION KEY(SectorSeq)
|
||||
REFERENCES
|
||||
Sector(SectorSeq) LABEL Belongs_To PROPERTIES ALL COLUMNS );
|
||||
@@ -0,0 +1,210 @@
|
||||
"""This file is for database operations done by the application"""
|
||||
|
||||
# pylint: disable=line-too-long
|
||||
import os
|
||||
|
||||
from dotenv import load_dotenv
|
||||
from google.api_core.client_options import ClientOptions
|
||||
from google.cloud import spanner
|
||||
import pandas as pd
|
||||
import streamlit as st
|
||||
from streamlit_extras.stylable_container import stylable_container
|
||||
|
||||
load_dotenv()
|
||||
|
||||
instance_id = os.getenv("instance_id")
|
||||
database_id = os.getenv("database_id")
|
||||
api_endpoint = os.getenv("api_endpoint")
|
||||
|
||||
options = ClientOptions(api_endpoint=api_endpoint)
|
||||
spanner_client = spanner.Client(client_options=options)
|
||||
|
||||
instance = spanner_client.instance(instance_id)
|
||||
database = instance.database(database_id)
|
||||
|
||||
|
||||
def spanner_read_data(query: str, *vector_input: list) -> pd.DataFrame:
|
||||
"""This function helps read data from Spanner"""
|
||||
with database.snapshot() as snapshot:
|
||||
if len(vector_input) != 0:
|
||||
results = snapshot.execute_sql(
|
||||
query,
|
||||
params={"vector": vector_input[0]},
|
||||
)
|
||||
else:
|
||||
results = snapshot.execute_sql(query)
|
||||
rows = list(results)
|
||||
cols = [x.name for x in results.fields]
|
||||
return pd.DataFrame(rows, columns=cols)
|
||||
|
||||
|
||||
def fts_query(query_params: list) -> dict:
|
||||
"""This function runs Full Text Search Query"""
|
||||
if query_params[1] == "":
|
||||
fts_query_str = (
|
||||
"SELECT DISTINCT fund_name,investment_strategy,investment_managers,fund_trailing_return_ytd,top5_holdings FROM EU_MutualFunds WHERE SEARCH(investment_strategy_Tokens, '"
|
||||
+ query_params[0]
|
||||
+ "') order by fund_name;"
|
||||
)
|
||||
else:
|
||||
fts_query_str = (
|
||||
"SELECT DISTINCT fund_name, manager, strategy, score FROM (SELECT fund_name , investment_managers AS manager, investment_strategy as strategy, SCORE_NGRAMS(investment_managers_Substring_Tokens_NGRAM, '"
|
||||
+ query_params[1]
|
||||
+ "') AS score FROM EU_MutualFunds WHERE SEARCH_NGRAMS(investment_managers_Substring_Tokens_NGRAM, '"
|
||||
+ query_params[1]
|
||||
+ "', min_ngrams=>1) AND SEARCH(investment_strategy_Tokens, '"
|
||||
+ query_params[0]
|
||||
+ "') ) ORDER BY score DESC;"
|
||||
)
|
||||
|
||||
return_vals = {}
|
||||
return_vals["query"] = fts_query_str
|
||||
df = spanner_read_data(fts_query_str)
|
||||
|
||||
return_vals["data"] = df
|
||||
return return_vals
|
||||
|
||||
|
||||
def semantic_query(query_params: list) -> dict:
|
||||
"""This function runs Semantic Text Search Query"""
|
||||
if query_params[1].strip() != "":
|
||||
semantic_query_string = (
|
||||
"SELECT fund_name, investment_strategy,investment_managers, COSINE_DISTANCE( investment_strategy_Embedding, (SELECT embeddings. VALUES FROM ML.PREDICT( MODEL EmbeddingsModel, (SELECT '"
|
||||
+ query_params[0]
|
||||
+ "' AS content) ) ) ) AS distance FROM EU_MutualFunds WHERE investment_strategy_Embedding is not NULL AND search_substring(investment_managers_substring_tokens, '"
|
||||
+ query_params[1]
|
||||
+ "')ORDER BY distance LIMIT 10;"
|
||||
)
|
||||
else:
|
||||
semantic_query_string = (
|
||||
"SELECT fund_name, investment_strategy,investment_managers, COSINE_DISTANCE( investment_strategy_Embedding, (SELECT embeddings. VALUES FROM ML.PREDICT( MODEL EmbeddingsModel, (SELECT '"
|
||||
+ query_params[0]
|
||||
+ "' AS content) ) ) ) AS distance FROM EU_MutualFunds WHERE investment_strategy_Embedding is not NULL ORDER BY distance LIMIT 10;"
|
||||
)
|
||||
return_vals = {}
|
||||
return_vals["query"] = semantic_query_string
|
||||
df = spanner_read_data(semantic_query_string)
|
||||
|
||||
return_vals["data"] = df
|
||||
return return_vals
|
||||
|
||||
|
||||
def semantic_query_ann(query_params: list) -> dict:
|
||||
"""This function runs Semantic Text Search ANN Query"""
|
||||
|
||||
embedding_query = (
|
||||
'SELECT embeddings. VALUES as vector FROM ML.PREDICT( MODEL EmbeddingsModel, (SELECT "'
|
||||
+ query_params[0]
|
||||
+ '" AS content) ) ;'
|
||||
)
|
||||
vector_input = spanner_read_data(embedding_query).values.tolist()
|
||||
|
||||
if query_params[1].strip() != "":
|
||||
ann_query = (
|
||||
"SELECT funds.fund_name, funds.investment_strategy, funds.investment_managers FROM (SELECT NewMFSequence, APPROX_EUCLIDEAN_DISTANCE(investment_strategy_Embedding_vector, @vector, options => JSON '{\"num_leaves_to_search\": 10}') AS distance FROM EU_MutualFunds @{force_index = InvestmentStrategyEmbeddingIndex} WHERE investment_strategy_Embedding_vector IS NOT NULL ORDER BY distance LIMIT 500 ) AS ann JOIN EU_MutualFunds AS funds ON ann.NewMFSequence = funds.NewMFSequence WHERE SEARCH_NGRAMS(funds.investment_managers_Substring_Tokens_NGRAM, '"
|
||||
+ query_params[1]
|
||||
+ "',min_ngrams=>1) ORDER BY SCORE_NGRAMS(funds.investment_managers_Substring_Tokens_NGRAM, '"
|
||||
+ query_params[1]
|
||||
+ "') desc;"
|
||||
)
|
||||
else:
|
||||
ann_query = "SELECT fund_name, investment_strategy, investment_managers, APPROX_EUCLIDEAN_DISTANCE(investment_strategy_Embedding_vector, @vector, options => JSON '{\"num_leaves_to_search\": 10}') AS distance FROM EU_MutualFunds @{force_index = InvestmentStrategyEmbeddingIndex} WHERE investment_strategy_Embedding_vector IS NOT NULL ORDER BY distance LIMIT 100;"
|
||||
results_df = spanner_read_data(ann_query, vector_input[0][0])
|
||||
results_df = spanner_read_data(ann_query, vector_input[0][0])
|
||||
results_df = spanner_read_data(ann_query, vector_input[0][0])
|
||||
results_df = spanner_read_data(ann_query, vector_input[0][0])
|
||||
|
||||
return_vals = {}
|
||||
return_vals["query"] = ann_query
|
||||
return_vals["data"] = results_df
|
||||
return return_vals
|
||||
|
||||
|
||||
def like_query(query_params: list) -> dict:
|
||||
"""This function runs Precise Text Search Query"""
|
||||
|
||||
if query_params[1] == "EXCLUDE":
|
||||
query_params[1] = "AND"
|
||||
precise_query = (
|
||||
" SELECT DISTINCT fund_name, investment_managers, investment_strategy FROM EU_MutualFunds WHERE investment_managers LIKE ('%"
|
||||
+ query_params[3]
|
||||
+ "%') AND ( investment_strategy LIKE ('%"
|
||||
+ query_params[0]
|
||||
+ "%') "
|
||||
+ query_params[1]
|
||||
+ " investment_strategy LIKE ('%"
|
||||
+ query_params[2]
|
||||
+ "%') ) ORDER BY fund_name;"
|
||||
)
|
||||
return_vals = {}
|
||||
return_vals["query"] = precise_query
|
||||
df = spanner_read_data(precise_query)
|
||||
|
||||
return_vals["data"] = df
|
||||
return return_vals
|
||||
|
||||
|
||||
def compliance_query(query_params: list) -> dict:
|
||||
"""This function runs Compliance Graph Search Query"""
|
||||
graph_compliance_query = (
|
||||
"GRAPH FundGraph MATCH (sector:Sector {sector_name: '"
|
||||
+ query_params[0]
|
||||
+ "'})<-[:BELONGS_TO]-(company:Company)<-[h:HOLDS]-(fund:Fund) RETURN fund.fund_name, SUM(h.percentage) AS totalHoldings GROUP BY fund.fund_name NEXT FILTER totalHoldings > "
|
||||
+ query_params[1]
|
||||
+ " RETURN fund_name, totalHoldings"
|
||||
)
|
||||
|
||||
return_vals = {}
|
||||
return_vals["query"] = graph_compliance_query
|
||||
df = spanner_read_data(graph_compliance_query)
|
||||
return_vals["data"] = df
|
||||
return return_vals
|
||||
|
||||
|
||||
def graph_dtls_query() -> dict:
|
||||
"""This function runs Graph Details Query"""
|
||||
company_query = "select CompanySeq,name from Companies;"
|
||||
|
||||
return_vals = {}
|
||||
df_companies = spanner_read_data(company_query)
|
||||
return_vals["Companies"] = df_companies
|
||||
|
||||
sector_query = "select * from Sectors;"
|
||||
df_sectors = spanner_read_data(sector_query)
|
||||
return_vals["Sectors"] = df_sectors
|
||||
|
||||
managers_query = "select * from Managers LIMIT 100;"
|
||||
df_managers = spanner_read_data(managers_query)
|
||||
return_vals["Managers"] = df_managers
|
||||
|
||||
company_belong_sector_query = "SELECT * from CompanyBelongsSector;"
|
||||
df_comp_sec_edge = spanner_read_data(company_belong_sector_query)
|
||||
return_vals["CompanySectorRelation"] = df_comp_sec_edge
|
||||
|
||||
mgr_fund_edge_query = " SELECT mgrs.NewMFSequence,fund_name,ManagerSeq from ManagerManagesFund mgrs JOIN EU_MutualFunds funds ON mgrs.NewMFSequence = funds.NewMFSequence where ManagerSeq in (select ManagerSeq from Managers LIMIT 100);"
|
||||
mgr_fund_edge = spanner_read_data(mgr_fund_edge_query)
|
||||
return_vals["ManagerFundRelation"] = mgr_fund_edge
|
||||
|
||||
funds_node_query = "select fund_name, NewMFSequence from EU_MutualFunds where NewMFSequence in (SELECT NewMFSequence FROM FundHoldsCompany);"
|
||||
funds_node = spanner_read_data(funds_node_query)
|
||||
return_vals["Funds"] = funds_node
|
||||
|
||||
funds_hold_company_edge_query = "SELECT * FROM FundHoldsCompany;"
|
||||
funds_hold_company_edge = spanner_read_data(funds_hold_company_edge_query)
|
||||
return_vals["FundsHoldsCompaniesRelation"] = funds_hold_company_edge
|
||||
|
||||
return return_vals
|
||||
|
||||
|
||||
def display_spanner_query(spanner_query: str) -> None:
|
||||
"""This function runs Graph Details Query"""
|
||||
with st.expander("Spanner Query"):
|
||||
with stylable_container(
|
||||
"codeblock",
|
||||
"""
|
||||
code {
|
||||
white-space: pre-wrap !important;
|
||||
}
|
||||
""",
|
||||
):
|
||||
st.code(spanner_query, language="sql", line_numbers=False)
|
||||
@@ -0,0 +1,50 @@
|
||||
"""This module is the page for Graph Viz Data Search feature"""
|
||||
|
||||
# pylint: disable=import-error, line-too-long, unused-variable
|
||||
|
||||
from database import graph_dtls_query
|
||||
from pyvis.network import Network
|
||||
|
||||
|
||||
def generate_graph() -> None:
|
||||
"""This function is for generating the Graph Visualization"""
|
||||
|
||||
graph = Network("900px", "900px", notebook=True, heading="")
|
||||
return_vals = graph_dtls_query()
|
||||
companies = return_vals.get("Companies")
|
||||
for index, row in companies.iterrows(): # type: ignore[union-attr] # might ignore other potential errors
|
||||
graph.add_node(
|
||||
str(row["CompanySeq"]),
|
||||
label=row["name"],
|
||||
title=row["name"],
|
||||
shape="triangle",
|
||||
)
|
||||
|
||||
sectors = return_vals.get("Sectors")
|
||||
for index, row in sectors.iterrows(): # type: ignore[union-attr] # might ignore other potential errors
|
||||
graph.add_node(
|
||||
str(row["SectorSeq"]),
|
||||
label=row["sector_name"],
|
||||
shape="square",
|
||||
color="red",
|
||||
title=row["sector_name"],
|
||||
)
|
||||
|
||||
funds = return_vals.get("Funds")
|
||||
for index, row in funds.iterrows(): # type: ignore[union-attr] # might ignore other potential errors
|
||||
graph.add_node(
|
||||
str(row["NewMFSequence"]),
|
||||
label=row["fund_name"],
|
||||
color="green",
|
||||
title=row["fund_name"],
|
||||
)
|
||||
|
||||
comp_sector_relation = return_vals.get("CompanySectorRelation")
|
||||
for index, row in comp_sector_relation.iterrows(): # type: ignore[union-attr] # might ignore other potential errors
|
||||
graph.add_edge(str(row["CompanySeq"]), str(row["SectorSeq"]), title="BELONGS")
|
||||
|
||||
fund_hold_company_relation = return_vals.get("FundsHoldsCompaniesRelation")
|
||||
for index, row in fund_hold_company_relation.iterrows(): # type: ignore[union-attr] # might ignore other potential errors
|
||||
graph.add_edge(str(row["NewMFSequence"]), str(row["CompanySeq"]), title="HOLDS")
|
||||
|
||||
graph.show("graph_viz.html")
|
||||
@@ -0,0 +1,35 @@
|
||||
"""This file is the home Page of the Python Streamlit app"""
|
||||
|
||||
# pylint: disable= import-error,line-too-long
|
||||
|
||||
import streamlit as st
|
||||
|
||||
st.set_page_config(
|
||||
layout="wide",
|
||||
page_title="FinVest Advisor",
|
||||
page_icon="https://storage.googleapis.com/github-repo/generative-ai/sample-apps/finance-advisor-spanner/images/small-logo.png",
|
||||
initial_sidebar_state="expanded",
|
||||
)
|
||||
|
||||
|
||||
st.logo(
|
||||
"https://storage.googleapis.com/github-repo/generative-ai/sample-apps/finance-advisor-spanner/images/investments.png"
|
||||
)
|
||||
|
||||
st.header("Welcome")
|
||||
st.image(
|
||||
"https://storage.googleapis.com/github-repo/generative-ai/sample-apps/finance-advisor-spanner/images/Finvest-white-removebg-preview.png"
|
||||
)
|
||||
|
||||
|
||||
def table_columns_layout_setup() -> dict:
|
||||
"""This function implements common layouts across the pages"""
|
||||
st.columns([0.25, 0.25, 0.20, 0.10])
|
||||
classes = ["display", "compact", "cell-border", "stripe"]
|
||||
buttons = ["pageLength", "csvHtml5", "excelHtml5", "colvis"]
|
||||
style = "table-layout:auto;width:auto;margin:auto;caption-side:bottom"
|
||||
it_args = {"classes": classes, "style": style}
|
||||
|
||||
if buttons:
|
||||
it_args["buttons"] = buttons
|
||||
return it_args
|
||||
@@ -0,0 +1,71 @@
|
||||
"""This module is the page for Asset Search feature"""
|
||||
|
||||
# pylint: disable=line-too-long,import-error,invalid-name
|
||||
|
||||
from database import display_spanner_query, fts_query, like_query
|
||||
from home import table_columns_layout_setup
|
||||
from itables.streamlit import interactive_table
|
||||
import streamlit as st
|
||||
|
||||
st.logo(
|
||||
"https://storage.googleapis.com/github-repo/generative-ai/sample-apps/finance-advisor-spanner/images/investments.png"
|
||||
)
|
||||
|
||||
|
||||
def asset_search_common(query_parameters: list, query_type: str) -> None:
|
||||
"""This function implements Asset search common functions"""
|
||||
|
||||
st.header("FinVest Fund Advisor")
|
||||
st.subheader("Asset Search")
|
||||
|
||||
with st.spinner("Querying Spanner..."):
|
||||
if query_type == "PRECISE":
|
||||
return_vals = like_query(query_parameters)
|
||||
else:
|
||||
return_vals = fts_query(query_parameters)
|
||||
spanner_query = return_vals.get("query")
|
||||
data = return_vals.get("data")
|
||||
display_spanner_query(str(spanner_query))
|
||||
|
||||
interactive_table(data, caption="", **table_columns_layout_setup())
|
||||
|
||||
|
||||
with st.sidebar:
|
||||
with st.form("Asset Search"):
|
||||
st.subheader("Search Criteria")
|
||||
precise_vs_text = st.radio("", ["Full-Text", "Precise"], horizontal=True)
|
||||
precise_search = False
|
||||
with st.expander("Asset Strategy", expanded=True):
|
||||
investment_strategy_pt1 = st.text_input("", value="Europe")
|
||||
and_or_exclude = st.radio("", ["AND", "OR", "EXCLUDE"], horizontal=True)
|
||||
investment_strategy_pt2 = st.text_input("", value="Asia")
|
||||
investment_manager = st.text_input("Investment Manager", value="James")
|
||||
investment_strategy = ""
|
||||
if precise_vs_text == "Full-Text":
|
||||
if and_or_exclude == "EXCLUDE":
|
||||
investment_strategy = (
|
||||
investment_strategy_pt1 + " -" + investment_strategy_pt2
|
||||
)
|
||||
else:
|
||||
investment_strategy = (
|
||||
investment_strategy_pt1
|
||||
+ " "
|
||||
+ and_or_exclude
|
||||
+ " "
|
||||
+ investment_strategy_pt2
|
||||
)
|
||||
else:
|
||||
precise_search = True
|
||||
asset_search_submitted = st.form_submit_button("Submit")
|
||||
if asset_search_submitted:
|
||||
if precise_search:
|
||||
query_params = [
|
||||
investment_strategy_pt1.strip(),
|
||||
and_or_exclude,
|
||||
investment_strategy_pt2.strip(),
|
||||
investment_manager.strip(),
|
||||
]
|
||||
asset_search_common(query_params, "PRECISE")
|
||||
else:
|
||||
query_params = [investment_strategy, investment_manager]
|
||||
asset_search_common(query_params, "FTS")
|
||||
@@ -0,0 +1,45 @@
|
||||
"""This module is the page for Semantic Search feature"""
|
||||
|
||||
# pylint: disable=line-too-long,import-error,invalid-name
|
||||
|
||||
from database import display_spanner_query, semantic_query, semantic_query_ann
|
||||
from home import table_columns_layout_setup
|
||||
from itables.streamlit import interactive_table
|
||||
import streamlit as st
|
||||
|
||||
st.logo(
|
||||
"https://storage.googleapis.com/github-repo/generative-ai/sample-apps/finance-advisor-spanner/images/investments.png"
|
||||
)
|
||||
|
||||
|
||||
def asset_semantic_search() -> None:
|
||||
"""This function implements Semantic Search feature"""
|
||||
|
||||
st.header("FinVest Fund Advisor")
|
||||
st.subheader("Semantic Search")
|
||||
query_params = [investment_strategy.strip(), investment_manager.strip()]
|
||||
|
||||
with st.spinner("Querying Spanner..."):
|
||||
if annVsKNN == "KNN":
|
||||
semantic_return_vals = semantic_query(query_params)
|
||||
else:
|
||||
semantic_return_vals = semantic_query_ann(query_params)
|
||||
semantic_queries = semantic_return_vals.get("query")
|
||||
data = semantic_return_vals.get("data")
|
||||
display_spanner_query(str(semantic_queries))
|
||||
|
||||
interactive_table(data, caption="", **table_columns_layout_setup())
|
||||
|
||||
|
||||
with st.sidebar:
|
||||
with st.form("Asset Semantic Search"):
|
||||
st.subheader("Search Criteria")
|
||||
annVsKNN = st.radio("", ["ANN", "KNN"], horizontal=True)
|
||||
investment_strategy = st.text_area(
|
||||
"Search for me",
|
||||
value="Invest in companies which also subscribe to my ideas around climate change, doing good for the planet",
|
||||
)
|
||||
investment_manager = st.text_input("Investment Manager", value="Maarten")
|
||||
asset_semantic_search_submitted = st.form_submit_button("Submit")
|
||||
if asset_semantic_search_submitted:
|
||||
asset_semantic_search()
|
||||
@@ -0,0 +1,24 @@
|
||||
"""This module is the page for Graph Visualization feature"""
|
||||
|
||||
# pylint: disable=line-too-long,import-error,invalid-name
|
||||
|
||||
import graph_viz
|
||||
import streamlit as st
|
||||
import streamlit.components.v1 as components
|
||||
|
||||
st.subheader("Show me the Relationships between Funds ,Companies and Sectors")
|
||||
|
||||
st.logo(
|
||||
"https://storage.googleapis.com/github-repo/generative-ai/sample-apps/finance-advisor-spanner/images/investments.png"
|
||||
)
|
||||
graph_viz.generate_graph()
|
||||
|
||||
with open("graph_viz.html", encoding="utf-8") as html_file:
|
||||
source_code = html_file.read()
|
||||
components.html(source_code, height=950, width=900)
|
||||
|
||||
with st.sidebar:
|
||||
st.subheader("Legend")
|
||||
st.image(
|
||||
"https://storage.googleapis.com/github-repo/generative-ai/sample-apps/finance-advisor-spanner/images/Graph-legend.png"
|
||||
)
|
||||
@@ -0,0 +1,48 @@
|
||||
"""This module is the page for Exposure Check Search feature"""
|
||||
|
||||
# pylint: disable=line-too-long,import-error,invalid-name
|
||||
|
||||
from database import compliance_query, display_spanner_query
|
||||
from home import table_columns_layout_setup
|
||||
from itables.streamlit import interactive_table
|
||||
import streamlit as st
|
||||
|
||||
st.logo(
|
||||
"https://storage.googleapis.com/github-repo/generative-ai/sample-apps/finance-advisor-spanner/images/investments.png"
|
||||
)
|
||||
|
||||
|
||||
def compliance_search() -> None:
|
||||
"""This function implements Compliance Check Graph feature"""
|
||||
st.header("FinVest Fund Advisor")
|
||||
st.subheader("Exposure Check")
|
||||
|
||||
query_params = []
|
||||
query_params.append(sectorOption)
|
||||
query_params.append(exposurePercentage)
|
||||
with st.spinner("Querying Spanner..."):
|
||||
compliance_vals = compliance_query(query_params)
|
||||
compliance_queries = compliance_vals.get("query")
|
||||
data = compliance_vals.get("data")
|
||||
display_spanner_query(str(compliance_queries))
|
||||
|
||||
interactive_table(data, caption="", **table_columns_layout_setup())
|
||||
|
||||
|
||||
with st.sidebar:
|
||||
with st.form("Compliance Search"):
|
||||
st.subheader("Search Criteria")
|
||||
sectorOption = st.selectbox(
|
||||
"Which sector would you want to focus on?",
|
||||
("Technology", "Pharma", "Semiconductors"),
|
||||
index=None,
|
||||
placeholder="Select sector ...",
|
||||
)
|
||||
exposurePercentage = st.select_slider(
|
||||
"How much exposure to this sector would you prefer",
|
||||
options=["10%", "20%", "30%", "40%", "50%", "60%", "70%"],
|
||||
)
|
||||
exposurePercentage = exposurePercentage[:2]
|
||||
compliance_search_submitted = st.form_submit_button("Submit")
|
||||
if compliance_search_submitted:
|
||||
compliance_search()
|
||||
@@ -0,0 +1,10 @@
|
||||
streamlit
|
||||
google-cloud-spanner
|
||||
itables==2.2.5
|
||||
streamlit-navigation-bar
|
||||
streamlit-extras
|
||||
streamlit-agraph
|
||||
SPARQLWrapper
|
||||
pyvis
|
||||
python-dotenv
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
outputs/
|
||||
@@ -0,0 +1,170 @@
|
||||
# gemini-hallcheck
|
||||
|
||||
Confidence-targeted, abstention-aware hallucination evaluator for Gemini via `google-genai`.
|
||||
|
||||
- Implements the Kalai et al. idea in practice: **"Answer only if you're > t confident; otherwise say `IDK`."**
|
||||
- Scores with an abstention-aware loss: **correct = +1, wrong = −t/(1−t), IDK = 0**.
|
||||
- Produces a **risk–coverage curve** (conditional accuracy vs. coverage) with **labeled t-points**.
|
||||
- Works from **CSV** or directly from **MMLU** (Hugging Face), with **random sampling** and an **`--idk-frac`** mixer to create **IDK-only** items (tests true abstention).
|
||||
- **LLM semantic judge** (Gemini 2.5 Flash-Lite) or **exact** judge.
|
||||
- **Async** with progress bar, **quota-aware retries** (honors server `RetryInfo`) and optional **client-side RPM throttle**.
|
||||
- Runs on **Gemini API** or **Vertex AI** (env-based switch).
|
||||
|
||||
> ℹ️ We do **not** implement MMLU-Pro here.
|
||||
|
||||
---
|
||||
|
||||
## Install
|
||||
|
||||
```bash
|
||||
python -m venv .venv && source .venv/bin/activate
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
### Auth options
|
||||
|
||||
**Gemini API (Developer API):**
|
||||
|
||||
```bash
|
||||
export GOOGLE_API_KEY=YOUR_KEY # or GEMINI_API_KEY
|
||||
```
|
||||
|
||||
**Vertex AI:**
|
||||
|
||||
```bash
|
||||
export GOOGLE_GENAI_USE_VERTEXAI=true
|
||||
export GOOGLE_CLOUD_PROJECT=your-gcp-project
|
||||
export GOOGLE_CLOUD_LOCATION=us-central1 # or europe-west1, etc.
|
||||
# Do NOT set GOOGLE_API_KEY when using Vertex AI mode.
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Quickstart
|
||||
|
||||
### CSV mode
|
||||
|
||||
```bash
|
||||
gemhall run --data examples/toy.csv --thresholds 0.5 0.75 0.9 --model gemini-2.5-flash-lite --progress --out outputs
|
||||
```
|
||||
|
||||
### MMLU (direct from HF Datasets)
|
||||
|
||||
```bash
|
||||
gemhall mmlu --thresholds 0.5 0.75 0.9 --model gemini-2.5-flash-lite --split test --subjects all --limit 200 --judge llm --async --concurrency 16 --progress --out outputs/mmlu
|
||||
```
|
||||
|
||||
**Mix in "IDK-only" items** (turns a fraction of sampled items into unanswerables so the only correct behavior is `IDK`):
|
||||
|
||||
```bash
|
||||
--idk-frac 0.3
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## What you get
|
||||
|
||||
- `results.csv` – one row per (item × t): prediction, abstained flag, correctness, score
|
||||
- `metrics.json` – coverage, conditional accuracy (on answers), hallucination rate among answers, avg expected score
|
||||
- `behavior.json` – simple behavior checks (e.g., monotonic coverage as t increases)
|
||||
- `rc_curve.png` – **risk–coverage curve** with **"t=…"** labels on each point
|
||||
- `report.md` – short summary plus the chart embedded
|
||||
|
||||
---
|
||||
|
||||
## Interpreting the curve
|
||||
|
||||
- **Coverage** (x-axis): fraction of items the model answered (didn't say `IDK`).
|
||||
- **Conditional accuracy** (y-axis): how often it was correct **when it did answer**.
|
||||
- As **t** rises ⇒ coverage falls, conditional accuracy should rise.
|
||||
- If **accuracy at t** is **well below t**, the model is **over-confident or non-compliant** → raise t, harden prompts, add retrieval/handoffs, or re-calibrate.
|
||||
|
||||
---
|
||||
|
||||
## CLI reference
|
||||
|
||||
Shared flags (both `run` and `mmlu`):
|
||||
|
||||
```sh
|
||||
--thresholds FLOAT... Confidence thresholds (e.g., 0.5 0.75 0.9) [required]
|
||||
--model TEXT Gemini model id (default: gemini-2.5-flash)
|
||||
--temperature FLOAT Sampling temperature (default: 0.0)
|
||||
--thinking-budget INT Optional thinking budget (default: 0)
|
||||
--seed INT RNG seed for sampling (default: 1234)
|
||||
--judge {exact,llm} Validity judge (exact or LLM; default: exact)
|
||||
--async Use async client for parallel requests
|
||||
--concurrency INT Max concurrent requests in async mode (default: 8)
|
||||
--progress Show a progress bar
|
||||
--out PATH Output directory (default: outputs)
|
||||
--rpm-limit INT Client-side requests-per-minute cap (optional)
|
||||
--max-retries INT Max retries on 429 with backoff (default: 6)
|
||||
```
|
||||
|
||||
`run` (CSV):
|
||||
|
||||
```sh
|
||||
--data PATH CSV with columns: id, question, gold, unknown_ok
|
||||
```
|
||||
|
||||
`mmlu` (Hugging Face "cais/mmlu"):
|
||||
|
||||
```sh
|
||||
--split TEXT Split (e.g., test, dev) [default: test]
|
||||
--subjects STR... Subject names or 'all' [default: all]
|
||||
--limit INT Randomly sample N items after filtering subjects
|
||||
--idk-frac FLOAT Fraction [0..1] to convert to IDK-only items (default: 0.0)
|
||||
```
|
||||
|
||||
**MMLU loader notes:**
|
||||
We first try the unified `"all"` config and fall back to stitching per-subject configs if needed. No `trust_remote_code` required. We also ensure a `subject` column exists.
|
||||
|
||||
---
|
||||
|
||||
## Judges
|
||||
|
||||
- **exact** – strict match for MCQ (letters A/B/C/D), or numerical/text equality for free-form.
|
||||
- **llm** – Gemini 2.5 Flash-Lite "YES/NO" grader; for `unknown_ok=1`, only `IDK` is considered correct (no LLM call).
|
||||
|
||||
---
|
||||
|
||||
## IDK detection & scoring
|
||||
|
||||
- We normalize model outputs; `IDK` is recognized case-insensitively with common variants.
|
||||
- Score per item at threshold **t**:
|
||||
- answered & correct: **+1**
|
||||
- answered & wrong: **−t/(1−t)**
|
||||
- abstained (`IDK`): **0**
|
||||
|
||||
---
|
||||
|
||||
## Rate limits & retries
|
||||
|
||||
- If you omit `--rpm-limit`, we still **auto-retry** on `429 RESOURCE_EXHAUSTED`, honoring server **`RetryInfo`** with jittered exponential backoff.
|
||||
- Set `--rpm-limit` to smooth out bursts and avoid 429s when running with high `--concurrency`.
|
||||
- Typical stable settings: `--async --concurrency 12 --rpm-limit 180 --max-retries 8`.
|
||||
|
||||
---
|
||||
|
||||
## Business mapping
|
||||
|
||||
- **t≈0.5: Drafting & triage** – high coverage, human-in-the-loop
|
||||
- **t≈0.75: Assistive answers** – support suggestions, FAQ with citations
|
||||
- **t≈0.9: Self-serve replies** – public answers in non-regulated flows
|
||||
- **t≈0.95: High-stakes** – regulated/brand-critical, else handoff
|
||||
|
||||
---
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
- **MMLU config error**: we now request `"all"` and fall back per-subject; ensure `datasets>=2.18.0`.
|
||||
- **Curly-brace crash in prompts**: fixed by using f-strings (brace-safe).
|
||||
- **429 quota**: use `--rpm-limit`, and/or lower `--concurrency`.
|
||||
- **Vertex AI vs API key**: set `GOOGLE_GENAI_USE_VERTEXAI=true` (+ project/location) to use Vertex AI; don't set an API key at the same time.
|
||||
|
||||
---
|
||||
|
||||
## Citation
|
||||
|
||||
This project is based on the following paper.
|
||||
|
||||
[arXiv:2509.04664](https://arxiv.org/abs/2509.04664)
|
||||
@@ -0,0 +1,4 @@
|
||||
id,question,gold,unknown_ok,category
|
||||
1,How many Ds are in LADDER? Reply with a number only.,2,0,count_letters
|
||||
2,What is the chemical symbol for gold?,Au,0,common_facts
|
||||
3,"In the Python 'requests' library, what does the `enable_turbo_mode()` function do?",,1,api_nonexistent
|
||||
|
@@ -0,0 +1,22 @@
|
||||
[project]
|
||||
name = "gemini-hallcheck"
|
||||
version = "0.4.0"
|
||||
description = "Confidence-targeted, abstention-aware hallucination evaluator for Gemini via google-genai"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
authors = [{name="Your Name or Org"}]
|
||||
dependencies = [
|
||||
"google-genai>=1.30.0",
|
||||
"pandas>=2.0.0",
|
||||
"matplotlib>=3.8.0",
|
||||
"numpy>=1.24.0",
|
||||
"tqdm>=4.66.0",
|
||||
"datasets>=2.18.0",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
gemhall = "gemhall.cli:main"
|
||||
|
||||
[build-system]
|
||||
requires = ["setuptools"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
Binary file not shown.
@@ -0,0 +1,201 @@
|
||||
import csv
|
||||
import os
|
||||
import random
|
||||
from collections.abc import Iterable
|
||||
|
||||
from datasets import concatenate_datasets, load_dataset
|
||||
|
||||
SUBJECTS_ALL = [
|
||||
"abstract_algebra",
|
||||
"anatomy",
|
||||
"astronomy",
|
||||
"business_ethics",
|
||||
"clinical_knowledge",
|
||||
"college_biology",
|
||||
"college_chemistry",
|
||||
"college_computer_science",
|
||||
"college_mathematics",
|
||||
"college_medicine",
|
||||
"college_physics",
|
||||
"computer_security",
|
||||
"conceptual_physics",
|
||||
"econometrics",
|
||||
"electrical_engineering",
|
||||
"elementary_mathematics",
|
||||
"formal_logic",
|
||||
"global_facts",
|
||||
"high_school_biology",
|
||||
"high_school_chemistry",
|
||||
"high_school_computer_science",
|
||||
"high_school_european_history",
|
||||
"high_school_geography",
|
||||
"high_school_government_and_politics",
|
||||
"high_school_macroeconomics",
|
||||
"high_school_mathematics",
|
||||
"high_school_microeconomics",
|
||||
"high_school_physics",
|
||||
"high_school_psychology",
|
||||
"high_school_statistics",
|
||||
"high_school_us_history",
|
||||
"high_school_world_history",
|
||||
"human_aging",
|
||||
"human_sexuality",
|
||||
"international_law",
|
||||
"jurisprudence",
|
||||
"logical_fallacies",
|
||||
"machine_learning",
|
||||
"management",
|
||||
"marketing",
|
||||
"medical_genetics",
|
||||
"miscellaneous",
|
||||
"moral_disputes",
|
||||
"moral_scenarios",
|
||||
"nutrition",
|
||||
"philosophy",
|
||||
"prehistory",
|
||||
"professional_accounting",
|
||||
"professional_law",
|
||||
"professional_medicine",
|
||||
"professional_psychology",
|
||||
"public_relations",
|
||||
"security_studies",
|
||||
"sociology",
|
||||
"us_foreign_policy",
|
||||
"virology",
|
||||
"world_religions",
|
||||
]
|
||||
|
||||
LETTERS = "ABCD"
|
||||
|
||||
|
||||
def _format_item(q: str, choices: list[str]) -> str:
|
||||
options = "\n".join(f"{LETTERS[i]}. {c}" for i, c in enumerate(choices[:4]))
|
||||
return f"{q}\n\nOptions:\n{options}\n\nReply with A, B, C, or D."
|
||||
|
||||
|
||||
def _gold_letter(answer_field) -> str:
|
||||
if isinstance(answer_field, int):
|
||||
return LETTERS[answer_field]
|
||||
s = str(answer_field).strip().upper()
|
||||
if s in set(LETTERS):
|
||||
return s
|
||||
try:
|
||||
return LETTERS[int(s)]
|
||||
except (ValueError, IndexError):
|
||||
raise ValueError(f"Invalid answer format: {answer_field}")
|
||||
|
||||
|
||||
def _answer_index(answer_field) -> int:
|
||||
if isinstance(answer_field, int):
|
||||
return int(answer_field)
|
||||
s = str(answer_field).strip().upper()
|
||||
if s in set(LETTERS):
|
||||
return LETTERS.index(s)
|
||||
return int(s)
|
||||
|
||||
|
||||
def export_temp_csv(
|
||||
out_csv: str,
|
||||
split: str = "test",
|
||||
subjects: Iterable[str] | None = None,
|
||||
limit: int | None = None,
|
||||
seed: int = 1234,
|
||||
idk_frac: float = 0.0,
|
||||
) -> str:
|
||||
"""Loads cais/mmlu from Hugging Face, filters subjects/split, optionally samples `limit` items,
|
||||
optionally converts a fraction into IDK-only, and writes a CSV for the evaluator.
|
||||
"""
|
||||
wanted_subjects = list(subjects) if subjects else SUBJECTS_ALL
|
||||
|
||||
# Try unified "all" config first
|
||||
table = None
|
||||
try:
|
||||
ds = load_dataset("cais/mmlu", name="all")
|
||||
if split not in ds:
|
||||
raise ValueError(
|
||||
f"Split '{split}' not found in cais/mmlu. Available: {list(ds.keys())}"
|
||||
)
|
||||
table = ds[split]
|
||||
if wanted_subjects != ["all"]:
|
||||
allowed = set(wanted_subjects)
|
||||
if "subject" in table.column_names:
|
||||
table = table.filter(lambda ex: ex.get("subject", "") in allowed)
|
||||
else:
|
||||
raise RuntimeError("Unified config missing 'subject' column")
|
||||
except Exception:
|
||||
# Fallback: concat selected subjects
|
||||
subjects_to_load = (
|
||||
SUBJECTS_ALL if wanted_subjects == ["all"] else wanted_subjects
|
||||
)
|
||||
parts = []
|
||||
for subj in subjects_to_load:
|
||||
d = load_dataset("cais/mmlu", name=subj)
|
||||
if split not in d:
|
||||
continue
|
||||
ds_split = d[split]
|
||||
if "subject" not in ds_split.column_names:
|
||||
ds_split = ds_split.map(lambda ex: {"subject": subj})
|
||||
parts.append(ds_split)
|
||||
if not parts:
|
||||
raise ValueError(
|
||||
f"No data found for split '{split}' and subjects {subjects_to_load}"
|
||||
)
|
||||
table = concatenate_datasets(parts)
|
||||
|
||||
# Random sampling after subject filtering
|
||||
if limit is not None and limit < len(table):
|
||||
rnd = random.Random(seed)
|
||||
idxs = rnd.sample(range(len(table)), limit)
|
||||
table = table.select(idxs)
|
||||
|
||||
# Decide which rows become IDK-only
|
||||
rnd = random.Random(seed)
|
||||
n = len(table)
|
||||
k = round(max(0.0, min(1.0, idk_frac)) * n)
|
||||
idk_idxs = set(rnd.sample(range(n), k)) if k > 0 else set()
|
||||
|
||||
rows: list[dict[str, str]] = []
|
||||
for i, ex in enumerate(table):
|
||||
q = ex["question"]
|
||||
choices = list(ex["choices"])
|
||||
if not isinstance(choices, list) or len(choices) < 4:
|
||||
continue
|
||||
subj = ex.get("subject", "mmlu")
|
||||
if i in idk_idxs:
|
||||
gi = _answer_index(ex["answer"]) # correct option index 0..3
|
||||
distractors = [j for j in range(4) if j != gi]
|
||||
repl = rnd.choice(distractors)
|
||||
choices[gi] = choices[
|
||||
repl
|
||||
] # duplicate a distractor -> no true option remains
|
||||
question = _format_item(q, choices[:4])
|
||||
rows.append(
|
||||
{
|
||||
"id": f"{subj}:{i}",
|
||||
"question": question,
|
||||
"gold": "",
|
||||
"unknown_ok": 1,
|
||||
"category": f"{subj}|unanswerable",
|
||||
}
|
||||
)
|
||||
else:
|
||||
ans = _gold_letter(ex["answer"])
|
||||
question = _format_item(q, choices[:4])
|
||||
rows.append(
|
||||
{
|
||||
"id": f"{subj}:{i}",
|
||||
"question": question,
|
||||
"gold": ans,
|
||||
"unknown_ok": 0,
|
||||
"category": subj,
|
||||
}
|
||||
)
|
||||
|
||||
os.makedirs(os.path.dirname(out_csv) or ".", exist_ok=True)
|
||||
with open(out_csv, "w", newline="", encoding="utf-8") as f:
|
||||
w = csv.DictWriter(
|
||||
f, fieldnames=["id", "question", "gold", "unknown_ok", "category"]
|
||||
)
|
||||
w.writeheader()
|
||||
w.writerows(rows)
|
||||
return out_csv
|
||||
@@ -0,0 +1,161 @@
|
||||
import argparse
|
||||
import contextlib
|
||||
import json
|
||||
import os
|
||||
from tempfile import NamedTemporaryFile
|
||||
|
||||
from gemhall.adapters.mmlu import SUBJECTS_ALL
|
||||
from gemhall.adapters.mmlu import export_temp_csv as mmlu_export
|
||||
from gemhall.eval import evaluate
|
||||
|
||||
|
||||
def add_common_args(ap: argparse.ArgumentParser) -> None:
|
||||
ap.add_argument(
|
||||
"--thresholds",
|
||||
nargs="+",
|
||||
type=float,
|
||||
required=True,
|
||||
help="Confidence thresholds, e.g., 0.5 0.75 0.9",
|
||||
)
|
||||
ap.add_argument("--out", default="outputs", help="Output directory")
|
||||
ap.add_argument(
|
||||
"--model",
|
||||
default="gemini-2.5-flash",
|
||||
choices=["gemini-2.0-flash", "gemini-2.5-flash", "gemini-2.5-flash-lite"],
|
||||
help="Gemini model id",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--temperature", type=float, default=0.0, help="Sampling temperature"
|
||||
)
|
||||
ap.add_argument(
|
||||
"--thinking-budget",
|
||||
type=int,
|
||||
default=0,
|
||||
dest="thinking_budget",
|
||||
help="Thinking budget (0 disables) - not supported at the moment",
|
||||
)
|
||||
ap.add_argument("--seed", type=int, default=1234, help="Sampling seed")
|
||||
ap.add_argument(
|
||||
"--judge",
|
||||
choices=["exact", "llm"],
|
||||
default="exact",
|
||||
help="Validity judge to use",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--async",
|
||||
action="store_true",
|
||||
dest="use_async",
|
||||
help="Use async client for parallelism",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--concurrency",
|
||||
type=int,
|
||||
default=8,
|
||||
help="Max concurrent requests in async mode",
|
||||
)
|
||||
ap.add_argument("--progress", action="store_true", help="Show a progress bar")
|
||||
ap.add_argument(
|
||||
"--rpm-limit",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Optional requests-per-minute cap (client-side)",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--max-retries",
|
||||
type=int,
|
||||
default=6,
|
||||
help="Max retries on 429/RESOURCE_EXHAUSTED",
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(
|
||||
prog="gemhall", description="Gemini confidence-targeted hallucination evaluator"
|
||||
)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
|
||||
# CSV mode
|
||||
run = sub.add_parser("run", help="Run evaluation from CSV")
|
||||
run.add_argument(
|
||||
"--data", required=True, help="Path to CSV with id,question,gold,unknown_ok"
|
||||
)
|
||||
add_common_args(run)
|
||||
|
||||
# MMLU mode (direct from HF)
|
||||
mmlu = sub.add_parser(
|
||||
"mmlu", help="Run evaluation on MMLU (direct from Hugging Face Datasets)"
|
||||
)
|
||||
mmlu.add_argument("--split", default="test", help="Dataset split (e.g., test, dev)")
|
||||
mmlu.add_argument(
|
||||
"--subjects", nargs="+", default=["all"], help="Subjects or 'all'"
|
||||
)
|
||||
mmlu.add_argument(
|
||||
"--limit",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Randomly sample N items after filtering subjects",
|
||||
)
|
||||
mmlu.add_argument(
|
||||
"--idk-frac",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Fraction [0..1] of sampled items to convert to IDK-only",
|
||||
)
|
||||
add_common_args(mmlu)
|
||||
|
||||
args = p.parse_args()
|
||||
|
||||
if args.cmd == "mmlu":
|
||||
with NamedTemporaryFile("w+", suffix=".csv", delete=False) as tf:
|
||||
tmp_csv = tf.name
|
||||
subjects = SUBJECTS_ALL if args.subjects == ["all"] else args.subjects
|
||||
csv_path = mmlu_export(
|
||||
out_csv=tmp_csv,
|
||||
split=args.split,
|
||||
subjects=subjects,
|
||||
limit=args.limit,
|
||||
seed=args.seed,
|
||||
idk_frac=args.idk_frac,
|
||||
)
|
||||
res = evaluate(
|
||||
data_csv=csv_path,
|
||||
thresholds=args.thresholds,
|
||||
out_dir=args.out,
|
||||
model=args.model,
|
||||
temperature=args.temperature,
|
||||
thinking_budget=0,
|
||||
seed=args.seed,
|
||||
judge=args.judge,
|
||||
use_async=args.use_async,
|
||||
concurrency=args.concurrency,
|
||||
show_progress=args.progress,
|
||||
rpm_limit=args.rpm_limit,
|
||||
max_retries=args.max_retries,
|
||||
)
|
||||
print(json.dumps(res, indent=2))
|
||||
with contextlib.suppress(Exception):
|
||||
os.unlink(csv_path)
|
||||
return
|
||||
|
||||
if args.cmd == "run":
|
||||
res = evaluate(
|
||||
data_csv=args.data,
|
||||
thresholds=args.thresholds,
|
||||
out_dir=args.out,
|
||||
model=args.model,
|
||||
temperature=args.temperature,
|
||||
thinking_budget=0,
|
||||
seed=args.seed,
|
||||
judge=args.judge,
|
||||
use_async=args.use_async,
|
||||
concurrency=args.concurrency,
|
||||
show_progress=args.progress,
|
||||
rpm_limit=args.rpm_limit,
|
||||
max_retries=args.max_retries,
|
||||
)
|
||||
print(json.dumps(res, indent=2))
|
||||
return
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,297 @@
|
||||
import asyncio
|
||||
import csv
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from gemhall.judge import judge_validity as judge_validity_exact
|
||||
from gemhall.judge_llm import AsyncLLMJudge, LLMJudge
|
||||
from gemhall.metrics import Record, aggregate, behavior_checks, score_item
|
||||
from gemhall.prompts import build_conf_prompt, is_idk
|
||||
from gemhall.runner import AsyncGeminiRunner, GeminiRunner
|
||||
|
||||
|
||||
def load_data_csv(path: str) -> list[dict[str, Any]]:
|
||||
rows: list[dict[str, Any]] = []
|
||||
with open(path, newline="", encoding="utf-8") as f:
|
||||
for row in csv.DictReader(f):
|
||||
rows.append(row)
|
||||
return rows
|
||||
|
||||
|
||||
def _write_artifacts(records: list[Record], out_dir: str) -> dict[str, str]:
|
||||
os.makedirs(out_dir, exist_ok=True)
|
||||
# results.csv
|
||||
results_csv = os.path.join(out_dir, "results.csv")
|
||||
with open(results_csv, "w", newline="", encoding="utf-8") as f:
|
||||
w = csv.writer(f)
|
||||
w.writerow(
|
||||
[
|
||||
"id",
|
||||
"t",
|
||||
"question",
|
||||
"gold",
|
||||
"unknown_ok",
|
||||
"pred",
|
||||
"abstained",
|
||||
"correct",
|
||||
"score",
|
||||
]
|
||||
)
|
||||
for r in records:
|
||||
w.writerow(
|
||||
[
|
||||
r.id,
|
||||
r.t,
|
||||
r.question,
|
||||
r.gold,
|
||||
int(r.unknown_ok),
|
||||
r.pred,
|
||||
int(r.abstained),
|
||||
int(r.correct),
|
||||
r.score,
|
||||
]
|
||||
)
|
||||
|
||||
# metrics/behavior
|
||||
metrics = aggregate(records)
|
||||
metrics_json = os.path.join(out_dir, "metrics.json")
|
||||
with open(metrics_json, "w", encoding="utf-8") as f:
|
||||
json.dump(metrics, f, indent=2)
|
||||
behavior = behavior_checks(metrics)
|
||||
behavior_json = os.path.join(out_dir, "behavior.json")
|
||||
with open(behavior_json, "w", encoding="utf-8") as f:
|
||||
json.dump(behavior, f, indent=2)
|
||||
|
||||
# plot with labels
|
||||
out_png = ""
|
||||
try:
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
ts = sorted(metrics.keys())
|
||||
cov = [metrics[t]["coverage"] for t in ts]
|
||||
acc = [metrics[t]["accuracy_conditioned_on_answering"] for t in ts]
|
||||
plt.figure()
|
||||
plt.plot(cov, acc, marker="o")
|
||||
for x, y, tval in zip(cov, acc, ts, strict=False):
|
||||
try:
|
||||
label = f"t={tval:g}"
|
||||
except Exception:
|
||||
label = f"t={tval}"
|
||||
plt.annotate(
|
||||
label, (x, y), textcoords="offset points", xytext=(6, 6), ha="left"
|
||||
)
|
||||
plt.xlabel("Coverage (answered fraction)")
|
||||
plt.ylabel("Conditional accuracy")
|
||||
plt.title("Risk–coverage curve (higher is better)")
|
||||
out_png = os.path.join(out_dir, "rc_curve.png")
|
||||
plt.savefig(out_png, bbox_inches="tight")
|
||||
plt.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# report.md
|
||||
report_md = os.path.join(out_dir, "report.md")
|
||||
with open(report_md, "w", encoding="utf-8") as f:
|
||||
f.write("# Hallucination Evaluation Report\n\n")
|
||||
f.write("## Summary\n")
|
||||
for t in sorted(metrics.keys()):
|
||||
m = metrics[t]
|
||||
f.write(
|
||||
f"- t={t}: coverage={m['coverage']:.3f}, cond_acc={m['accuracy_conditioned_on_answering']:.3f}, halluc_rate={m['hallucination_rate_among_answers']:.3f}, avg_score={m['avg_expected_score']:.3f}\n"
|
||||
)
|
||||
f.write("\n## Behavioral Check\n")
|
||||
f.write(json.dumps(behavior, indent=2))
|
||||
if out_png:
|
||||
f.write("\n\n## Risk–Coverage Curve\n\n")
|
||||
f.write(f"})\n")
|
||||
|
||||
return {
|
||||
"results_csv": results_csv,
|
||||
"metrics_json": metrics_json,
|
||||
"behavior_json": behavior_json,
|
||||
"report_md": report_md,
|
||||
"rc_curve_png": out_png or "",
|
||||
}
|
||||
|
||||
|
||||
def evaluate_sync(
|
||||
data_csv: str,
|
||||
thresholds: Sequence[float],
|
||||
out_dir: str,
|
||||
model: str,
|
||||
temperature: float,
|
||||
thinking_budget: int,
|
||||
seed: int,
|
||||
judge: str = "exact",
|
||||
show_progress: bool = False,
|
||||
rpm_limit: int | None = None,
|
||||
max_retries: int = 6,
|
||||
) -> dict[str, Any]:
|
||||
rows = load_data_csv(data_csv)
|
||||
runner = GeminiRunner(
|
||||
model=model,
|
||||
temperature=temperature,
|
||||
thinking_budget=thinking_budget,
|
||||
seed=seed,
|
||||
rpm_limit=rpm_limit,
|
||||
max_retries=max_retries,
|
||||
)
|
||||
llm_judge = LLMJudge() if judge == "llm" else None
|
||||
|
||||
total = len(rows) * len(thresholds)
|
||||
pbar = tqdm(total=total, disable=not show_progress, desc="Evaluating (sync)")
|
||||
|
||||
records: list[Record] = []
|
||||
for row in rows:
|
||||
qid = row.get("id") or ""
|
||||
question = row.get("question") or ""
|
||||
gold = row.get("gold") or ""
|
||||
unknown_ok = str(row.get("unknown_ok") or "0").strip() in {"1", "true", "True"}
|
||||
for t in thresholds:
|
||||
prompt = build_conf_prompt(question, float(t))
|
||||
pred = runner.generate(prompt)
|
||||
abstained = is_idk(pred)
|
||||
correct = (
|
||||
llm_judge.judge(question, gold, pred, unknown_ok)
|
||||
if judge == "llm"
|
||||
else judge_validity_exact(pred, gold, unknown_ok)
|
||||
)
|
||||
score = score_item(answered=not abstained, correct=correct, t=float(t))
|
||||
records.append(
|
||||
Record(
|
||||
id=qid,
|
||||
t=float(t),
|
||||
question=question,
|
||||
gold=gold,
|
||||
unknown_ok=unknown_ok,
|
||||
pred=pred,
|
||||
abstained=abstained,
|
||||
correct=correct,
|
||||
score=score,
|
||||
)
|
||||
)
|
||||
pbar.update(1)
|
||||
pbar.close()
|
||||
return _write_artifacts(records, out_dir)
|
||||
|
||||
|
||||
async def evaluate_async(
|
||||
data_csv: str,
|
||||
thresholds: Sequence[float],
|
||||
out_dir: str,
|
||||
model: str,
|
||||
temperature: float,
|
||||
thinking_budget: int,
|
||||
seed: int,
|
||||
judge: str = "exact",
|
||||
concurrency: int = 8,
|
||||
show_progress: bool = False,
|
||||
rpm_limit: int | None = None,
|
||||
max_retries: int = 6,
|
||||
) -> dict[str, Any]:
|
||||
rows = load_data_csv(data_csv)
|
||||
runner = AsyncGeminiRunner(
|
||||
model=model,
|
||||
temperature=temperature,
|
||||
thinking_budget=thinking_budget,
|
||||
seed=seed,
|
||||
rpm_limit=rpm_limit,
|
||||
max_retries=max_retries,
|
||||
)
|
||||
llm_judge = AsyncLLMJudge() if judge == "llm" else None
|
||||
|
||||
import asyncio
|
||||
|
||||
sem = asyncio.Semaphore(concurrency)
|
||||
total = len(rows) * len(thresholds)
|
||||
pbar = tqdm(total=total, disable=not show_progress, desc="Evaluating (async)")
|
||||
records: list[Record] = []
|
||||
|
||||
async def one(qid, question, gold, unknown_ok, t):
|
||||
prompt = build_conf_prompt(question, float(t))
|
||||
async with sem:
|
||||
pred = await runner.generate(prompt)
|
||||
abstained = is_idk(pred)
|
||||
if judge == "llm":
|
||||
async with sem:
|
||||
correct = await llm_judge.judge(question, gold, pred, unknown_ok)
|
||||
else:
|
||||
correct = judge_validity_exact(pred, gold, unknown_ok)
|
||||
score = score_item(answered=not abstained, correct=correct, t=float(t))
|
||||
return Record(
|
||||
id=qid,
|
||||
t=float(t),
|
||||
question=question,
|
||||
gold=gold,
|
||||
unknown_ok=unknown_ok,
|
||||
pred=pred,
|
||||
abstained=abstained,
|
||||
correct=correct,
|
||||
score=score,
|
||||
)
|
||||
|
||||
tasks = []
|
||||
for row in rows:
|
||||
qid = row.get("id") or ""
|
||||
question = row.get("question") or ""
|
||||
gold = row.get("gold") or ""
|
||||
unknown_ok = str(row.get("unknown_ok") or "0").strip() in {"1", "true", "True"}
|
||||
for t in thresholds:
|
||||
tasks.append(one(qid, question, gold, unknown_ok, float(t)))
|
||||
|
||||
for coro in asyncio.as_completed(tasks):
|
||||
records.append(await coro)
|
||||
pbar.update(1)
|
||||
pbar.close()
|
||||
return _write_artifacts(records, out_dir)
|
||||
|
||||
|
||||
def evaluate(
|
||||
data_csv: str,
|
||||
thresholds: Sequence[float],
|
||||
out_dir: str,
|
||||
model: str = "gemini-2.5-flash",
|
||||
temperature: float = 0.0,
|
||||
thinking_budget: int = 0,
|
||||
seed: int = 1234,
|
||||
judge: str = "exact",
|
||||
use_async: bool = False,
|
||||
concurrency: int = 8,
|
||||
show_progress: bool = False,
|
||||
rpm_limit: int | None = None,
|
||||
max_retries: int = 6,
|
||||
) -> dict[str, Any]:
|
||||
if use_async:
|
||||
return asyncio.run(
|
||||
evaluate_async(
|
||||
data_csv,
|
||||
thresholds,
|
||||
out_dir,
|
||||
model,
|
||||
temperature,
|
||||
thinking_budget,
|
||||
seed,
|
||||
judge,
|
||||
concurrency,
|
||||
show_progress,
|
||||
rpm_limit,
|
||||
max_retries,
|
||||
)
|
||||
)
|
||||
return evaluate_sync(
|
||||
data_csv,
|
||||
thresholds,
|
||||
out_dir,
|
||||
model,
|
||||
temperature,
|
||||
thinking_budget,
|
||||
seed,
|
||||
judge,
|
||||
show_progress,
|
||||
rpm_limit,
|
||||
max_retries,
|
||||
)
|
||||
@@ -0,0 +1,38 @@
|
||||
import re
|
||||
from collections.abc import Callable
|
||||
|
||||
|
||||
def normalize(s: str | None) -> str:
|
||||
if s is None:
|
||||
return ""
|
||||
s = str(s).strip().lower()
|
||||
s = re.sub(r"\s+", " ", s)
|
||||
return re.sub(r"[\s\.,;:!\?\-_/\\]+$", "", s)
|
||||
|
||||
|
||||
def numbers_equal(a: str, b: str) -> bool:
|
||||
try:
|
||||
return float(a) == float(b)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def str_or_num_equal(a: str, b: str) -> bool:
|
||||
A, B = normalize(a), normalize(b)
|
||||
if A == B:
|
||||
return True
|
||||
return numbers_equal(A, B)
|
||||
|
||||
|
||||
def judge_validity(
|
||||
pred: str,
|
||||
gold: str,
|
||||
unknown_ok: bool,
|
||||
eq_fn: Callable[[str, str], bool] | None = None,
|
||||
) -> bool:
|
||||
s = normalize(pred)
|
||||
if unknown_ok:
|
||||
return s == "idk"
|
||||
if eq_fn is None:
|
||||
eq_fn = str_or_num_equal
|
||||
return bool(eq_fn(pred, gold))
|
||||
@@ -0,0 +1,80 @@
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
|
||||
LLM_JUDGE_MODEL = "gemini-2.5-flash-lite"
|
||||
|
||||
SYSTEM = (
|
||||
"You are a strict binary grader. Decide if the candidate answer is semantically "
|
||||
"equivalent to the gold answer for the user question. "
|
||||
"Output ONLY one token: 'YES' or 'NO'. "
|
||||
"Be strict about factual equivalence, units, and exactness when appropriate. "
|
||||
"If the candidate says 'IDK' and the gold is blank, treat as 'YES'."
|
||||
)
|
||||
|
||||
PROMPT_TMPL = (
|
||||
"Question: {question}\n"
|
||||
"Gold answer: {gold}\n"
|
||||
"Candidate answer: {pred}\n"
|
||||
"Respond strictly with YES or NO."
|
||||
)
|
||||
|
||||
|
||||
def _postprocess(text: str | None) -> str:
|
||||
s = (text or "").strip().upper()
|
||||
if "YES" in s and "NO" not in s:
|
||||
return "YES"
|
||||
if "NO" in s and "YES" not in s:
|
||||
return "NO"
|
||||
return "YES" if s.startswith("Y") else "NO"
|
||||
|
||||
|
||||
class LLMJudge:
|
||||
def __init__(
|
||||
self,
|
||||
model: str = LLM_JUDGE_MODEL,
|
||||
temperature: float = 0.0,
|
||||
seed: int | None = 1234,
|
||||
):
|
||||
self.client = genai.Client()
|
||||
self.model = model
|
||||
self.temperature = temperature
|
||||
self.seed = seed
|
||||
|
||||
def judge(self, question: str, gold: str, pred: str, unknown_ok: bool) -> bool:
|
||||
if unknown_ok:
|
||||
return pred.strip().upper() == "IDK"
|
||||
contents = f"Question: {question}\nGold answer: {gold}\nCandidate answer: {pred}\nRespond strictly with YES or NO."
|
||||
cfg = types.GenerateContentConfig(
|
||||
temperature=self.temperature, seed=self.seed, max_output_tokens=4
|
||||
)
|
||||
resp = self.client.models.generate_content(
|
||||
model=self.model, contents=f"{SYSTEM}\n\n{contents}", config=cfg
|
||||
)
|
||||
return _postprocess(resp.text) == "YES"
|
||||
|
||||
|
||||
class AsyncLLMJudge:
|
||||
def __init__(
|
||||
self,
|
||||
model: str = LLM_JUDGE_MODEL,
|
||||
temperature: float = 0.0,
|
||||
seed: int | None = 1234,
|
||||
):
|
||||
self.client = genai.Client()
|
||||
self.model = model
|
||||
self.temperature = temperature
|
||||
self.seed = seed
|
||||
|
||||
async def judge(
|
||||
self, question: str, gold: str, pred: str, unknown_ok: bool
|
||||
) -> bool:
|
||||
if unknown_ok:
|
||||
return pred.strip().upper() == "IDK"
|
||||
contents = PROMPT_TMPL.format(question=question, gold=gold, pred=pred)
|
||||
cfg = types.GenerateContentConfig(
|
||||
temperature=self.temperature, seed=self.seed, max_output_tokens=4
|
||||
)
|
||||
resp = await self.client.aio.models.generate_content(
|
||||
model=self.model, contents=f"{SYSTEM}\n\n{contents}", config=cfg
|
||||
)
|
||||
return _postprocess(resp.text) == "YES"
|
||||
@@ -0,0 +1,66 @@
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class Record:
|
||||
id: str
|
||||
t: float
|
||||
question: str
|
||||
gold: str
|
||||
unknown_ok: bool
|
||||
pred: str
|
||||
abstained: bool
|
||||
correct: bool
|
||||
score: float
|
||||
|
||||
|
||||
def score_item(answered: bool, correct: bool, t: float) -> float:
|
||||
if not answered:
|
||||
return 0.0
|
||||
if t >= 1.0:
|
||||
return 1.0 if correct else (-1.0) * float("inf")
|
||||
return 1.0 if correct else -t / (1.0 - t)
|
||||
|
||||
|
||||
def aggregate(records: list[Record]) -> dict[float, dict[str, Any]]:
|
||||
by_t: dict[float, list[Record]] = defaultdict(list)
|
||||
for r in records:
|
||||
by_t[r.t].append(r)
|
||||
out: dict[float, dict[str, Any]] = {}
|
||||
for t, recs in by_t.items():
|
||||
n = len(recs)
|
||||
answered = sum(0 if r.abstained else 1 for r in recs)
|
||||
correct_ans = sum(1 for r in recs if (not r.abstained and r.correct))
|
||||
incorrect_ans = sum(1 for r in recs if (not r.abstained and not r.correct))
|
||||
coverage = answered / n if n else 0.0
|
||||
acc_cond = (correct_ans / answered) if answered else 0.0
|
||||
halluc_rate = (incorrect_ans / answered) if answered else 0.0
|
||||
total_score = sum(r.score for r in recs)
|
||||
out[t] = {
|
||||
"n": n,
|
||||
"coverage": coverage,
|
||||
"accuracy_conditioned_on_answering": acc_cond,
|
||||
"hallucination_rate_among_answers": halluc_rate,
|
||||
"avg_expected_score": total_score / n if n else 0.0,
|
||||
"answered": answered,
|
||||
"correct_answers": correct_ans,
|
||||
"incorrect_answers": incorrect_ans,
|
||||
"abstentions": n - answered,
|
||||
}
|
||||
return out
|
||||
|
||||
|
||||
def behavior_checks(metrics: dict[float, dict[str, Any]]) -> dict[str, Any]:
|
||||
ts = sorted(metrics.keys())
|
||||
covs = [metrics[t]["coverage"] for t in ts]
|
||||
monotone_violations = sum(
|
||||
1 for i in range(1, len(covs)) if covs[i] > covs[i - 1] + 1e-6
|
||||
)
|
||||
return {
|
||||
"thresholds": ts,
|
||||
"coverage": covs,
|
||||
"monotonic_coverage_expected": True,
|
||||
"monotonicity_violations": monotone_violations,
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
IDK_TOKENS = {
|
||||
"idk",
|
||||
"i don't know",
|
||||
"i do not know",
|
||||
"unknown",
|
||||
"cannot answer",
|
||||
"can't answer",
|
||||
"not sure",
|
||||
}
|
||||
|
||||
|
||||
def build_conf_prompt(question: str, t: float) -> str:
|
||||
pen = round(t / (1.0 - t), 3) if t < 1 else "infinite"
|
||||
t_disp = f"{t:.2f}".rstrip("0").rstrip(".")
|
||||
return (
|
||||
f"{question}\n\n"
|
||||
f"Answer only if you are > {t_disp} confident.\n"
|
||||
f"Scoring: correct = +1, incorrect = -{pen}, 'IDK' = 0.\n"
|
||||
f"If you are not > {t_disp} confident, reply EXACTLY 'IDK'.\n"
|
||||
"Output only the final answer (or 'IDK')."
|
||||
)
|
||||
|
||||
|
||||
def is_idk(text: str | None) -> bool:
|
||||
if text is None:
|
||||
return True
|
||||
s = text.strip().lower()
|
||||
return s == "idk" or s in IDK_TOKENS
|
||||
@@ -0,0 +1,180 @@
|
||||
import asyncio
|
||||
import random
|
||||
import time
|
||||
from collections import deque
|
||||
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
from google.genai.errors import ClientError
|
||||
|
||||
|
||||
class _AsyncRateLimiter:
|
||||
def __init__(self, rpm: int | None = None):
|
||||
self.rpm = rpm
|
||||
self._win = 60.0
|
||||
self._q = deque()
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
async def acquire(self) -> None:
|
||||
if not self.rpm:
|
||||
return
|
||||
while True:
|
||||
async with self._lock:
|
||||
now = asyncio.get_event_loop().time()
|
||||
while self._q and now - self._q[0] > self._win:
|
||||
self._q.popleft()
|
||||
if len(self._q) < self.rpm:
|
||||
self._q.append(now)
|
||||
return
|
||||
sleep_for = self._win - (now - self._q[0]) + 0.01
|
||||
await asyncio.sleep(sleep_for)
|
||||
|
||||
|
||||
class _RateLimiter:
|
||||
def __init__(self, rpm: int | None = None):
|
||||
self.rpm = rpm
|
||||
self._win = 60.0
|
||||
self._q = deque()
|
||||
|
||||
def acquire(self) -> None:
|
||||
if not self.rpm:
|
||||
return
|
||||
while True:
|
||||
now = time.monotonic()
|
||||
while self._q and now - self._q[0] > self._win:
|
||||
self._q.popleft()
|
||||
if len(self._q) < self.rpm:
|
||||
self._q.append(now)
|
||||
return
|
||||
sleep_for = self._win - (now - self._q[0]) + 0.01
|
||||
time.sleep(sleep_for)
|
||||
|
||||
|
||||
class GeminiRunner:
|
||||
def __init__(
|
||||
self,
|
||||
model: str = "gemini-2.5-flash",
|
||||
*,
|
||||
temperature: float = 0.0,
|
||||
thinking_budget: int = 0,
|
||||
seed: int | None = 1234,
|
||||
rpm_limit: int | None = None,
|
||||
max_retries: int = 6,
|
||||
):
|
||||
self.client = genai.Client()
|
||||
self.model = model
|
||||
self.temperature = temperature
|
||||
self.thinking_budget = thinking_budget
|
||||
self.seed = seed
|
||||
self.max_retries = max_retries
|
||||
self._rl = _RateLimiter(rpm_limit)
|
||||
|
||||
def generate(self, prompt: str) -> str:
|
||||
cfg = types.GenerateContentConfig(
|
||||
temperature=self.temperature,
|
||||
seed=self.seed,
|
||||
thinking_config=types.ThinkingConfig(thinking_budget=self.thinking_budget)
|
||||
if self.thinking_budget is not None
|
||||
else None,
|
||||
max_output_tokens=128,
|
||||
)
|
||||
backoff = 1.0
|
||||
for attempt in range(1, self.max_retries + 1):
|
||||
self._rl.acquire()
|
||||
try:
|
||||
resp = self.client.models.generate_content(
|
||||
model=self.model, contents=prompt, config=cfg
|
||||
)
|
||||
return (resp.text or "").strip()
|
||||
except ClientError as e:
|
||||
if (
|
||||
getattr(e, "status_code", None) == 429
|
||||
and attempt < self.max_retries
|
||||
):
|
||||
delay = None
|
||||
try:
|
||||
details = (
|
||||
(getattr(e, "response_json", {}) or {})
|
||||
.get("error", {})
|
||||
.get("details", [])
|
||||
)
|
||||
for d in details:
|
||||
if d.get("@type", "").endswith("RetryInfo"):
|
||||
s = d.get("retryDelay", "0s")
|
||||
if s.endswith("s"):
|
||||
delay = float(s[:-1])
|
||||
break
|
||||
except Exception:
|
||||
pass
|
||||
sleep_s = delay if delay is not None else min(60.0, backoff)
|
||||
sleep_s *= 0.8 + 0.4 * random.random() # jitter
|
||||
time.sleep(sleep_s)
|
||||
backoff *= 2
|
||||
continue
|
||||
raise
|
||||
return None
|
||||
|
||||
|
||||
class AsyncGeminiRunner:
|
||||
def __init__(
|
||||
self,
|
||||
model: str = "gemini-2.5-flash",
|
||||
*,
|
||||
temperature: float = 0.0,
|
||||
thinking_budget: int = 0,
|
||||
seed: int | None = 1234,
|
||||
rpm_limit: int | None = None,
|
||||
max_retries: int = 6,
|
||||
):
|
||||
self.client = genai.Client()
|
||||
self.model = model
|
||||
self.temperature = temperature
|
||||
self.thinking_budget = thinking_budget
|
||||
self.seed = seed
|
||||
self.max_retries = max_retries
|
||||
self._rl = _AsyncRateLimiter(rpm_limit)
|
||||
|
||||
async def generate(self, prompt: str) -> str:
|
||||
cfg = types.GenerateContentConfig(
|
||||
temperature=self.temperature,
|
||||
seed=self.seed,
|
||||
thinking_config=types.ThinkingConfig(thinking_budget=self.thinking_budget)
|
||||
if self.thinking_budget is not None
|
||||
else None,
|
||||
max_output_tokens=128,
|
||||
)
|
||||
backoff = 1.0
|
||||
for attempt in range(1, self.max_retries + 1):
|
||||
await self._rl.acquire()
|
||||
try:
|
||||
resp = await self.client.aio.models.generate_content(
|
||||
model=self.model, contents=prompt, config=cfg
|
||||
)
|
||||
return (resp.text or "").strip()
|
||||
except ClientError as e:
|
||||
if (
|
||||
getattr(e, "status_code", None) == 429
|
||||
and attempt < self.max_retries
|
||||
):
|
||||
delay = None
|
||||
try:
|
||||
details = (
|
||||
(getattr(e, "response_json", {}) or {})
|
||||
.get("error", {})
|
||||
.get("details", [])
|
||||
)
|
||||
for d in details:
|
||||
if d.get("@type", "").endswith("RetryInfo"):
|
||||
s = d.get("retryDelay", "0s")
|
||||
if s.endswith("s"):
|
||||
delay = float(s[:-1])
|
||||
break
|
||||
except Exception:
|
||||
pass
|
||||
sleep_s = delay if delay is not None else min(60.0, backoff)
|
||||
sleep_s *= 0.8 + 0.4 * random.random() # jitter
|
||||
await asyncio.sleep(sleep_s)
|
||||
backoff *= 2
|
||||
continue
|
||||
raise
|
||||
return None
|
||||
@@ -0,0 +1,25 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# https://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
#Google Cloud Configuration
|
||||
GOOGLE_CLOUD_PROJECT=your-google-cloud-project-id
|
||||
GOOGLE_CLOUD_LOCATION=us-central1
|
||||
GOOGLE_GENAI_USE_VERTEXAI=TRUE
|
||||
|
||||
# Gemini Model Configuration
|
||||
GOOGLE_GENAI_MODEL=gemini-live-2.5-flash-native-audio
|
||||
|
||||
# Update with your ngrok service forwarding URL for local testing only
|
||||
# Please remove this in cloud deployments
|
||||
SERVICE_URL="YOUR_NGROK_SERVICE_FORWARDING_URL" # e.g., https://<random-string>.ngrok-free.dev
|
||||
@@ -0,0 +1,40 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# https://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# Use a specific, slim Python version for reproducibility
|
||||
FROM python:3.12-slim
|
||||
|
||||
# Set the working directory in the container
|
||||
WORKDIR /app
|
||||
|
||||
# Install system dependencies required for audio processing and building packages
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
build-essential \
|
||||
cmake \
|
||||
git \
|
||||
libsamplerate0 \
|
||||
&& apt-get clean && rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Copy the requirements file and install Python dependencies
|
||||
COPY requirements.txt .
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
# Copy the rest of the application code into the container
|
||||
COPY . .
|
||||
|
||||
# Set the command to run the application using Gunicorn as the production server
|
||||
# Gunicorn manages Uvicorn workers to run the FastAPI app.
|
||||
# -b :$PORT binds to the port specified by the Cloud Run environment variable.
|
||||
# --workers 1 and --threads 8 is a good starting point for this CPU-bound app.
|
||||
CMD exec gunicorn --bind :$PORT --workers 1 --worker-class uvicorn.workers.UvicornWorker --threads 8 main:app
|
||||
@@ -0,0 +1,140 @@
|
||||
# AI Care Assistant
|
||||
|
||||
## Overview
|
||||
|
||||
This project provides a comprehensive architectural blueprint and implementation for a real-time, bidirectional voice-to-AI application. It integrates Twilio for telephony, a FastAPI backend for real-time processing, and the Google Gemini Live API for conversational AI. The application is designed for low-latency, high-fidelity conversational experiences, addressing challenges in system-level integration, audio transcoding, and deployment on Google Cloud Run.
|
||||
|
||||
Key design decisions include the use of `python-samplerate` for high-quality streaming audio resampling and a Cloud Run deployment strategy that mitigates cold starts (`min-instances=1`) and manages state using session affinity for in-memory DSP state.
|
||||
|
||||
<p align="center">
|
||||
<b></b> Click the image below to watch the video! </b>
|
||||
</p>
|
||||
<a href="https://www.youtube.com/watch?v=sboxrNY57uA">
|
||||
<img
|
||||
src="https://img.youtube.com/vi/sboxrNY57uA/maxresdefault.jpg"
|
||||
alt="Demo Video"
|
||||
style="width:50%; margin: 0 auto; display: block;"
|
||||
>
|
||||
</a>
|
||||
|
||||
## Special Notes
|
||||
|
||||
For a detailed understanding of the system's architecture, including component breakdowns, data flow, technical justifications, step-by-step guide on the implementation process, including project setup, code structure, and deployment instructions, please refer to the `design_doc.md` file.
|
||||
|
||||
## Google Cloud and Gemini Setup
|
||||
|
||||
1. **Set up a Google Cloud Project:**
|
||||
- Go to the [Google Cloud Console](https://console.cloud.google.com/) and create a new project.
|
||||
- Make sure to enable the Vertex AI API for your project.
|
||||
|
||||
2. **Authenticate with your Google Cloud Platform account:**
|
||||
- In your local terminal, authenticate your Google Cloud account by running:
|
||||
```bash
|
||||
gcloud auth login
|
||||
gcloud auth application-default login # for providing credentials to applications and code
|
||||
```
|
||||
3. **Set environment variables:**
|
||||
Set your environment variables by creating a **.env file** in the project directory by utilizing the **.env.example file**
|
||||
|
||||
## Quickstart (For local testing)
|
||||
|
||||
This project is implemented using Python 3.12.
|
||||
|
||||
1. **Create and activate a virtual environment:**
|
||||
```bash
|
||||
python3.12 -m venv venv
|
||||
source venv/bin/activate
|
||||
```
|
||||
|
||||
2. **Install the necessary packages:**
|
||||
```bash
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
3. **Install the ngrok:**
|
||||
To install ngrok on Linux, you can follow these steps:
|
||||
|
||||
- ***Download ngrok:***
|
||||
Open your web browser and go to the ngrok download page (https://ngrok.com/download). Download the Linux version.
|
||||
|
||||
- ***Unzip the file:***
|
||||
Open a terminal and navigate to your Downloads directory (or wherever you saved the file). Then unzip it:
|
||||
unzip /path/to/ngrok-v3-stable-linux-amd64.zip
|
||||
(Replace /path/to/ with the actual path to the downloaded file).
|
||||
|
||||
- ***Move ngrok to your PATH:***
|
||||
To make ngrok accessible from any directory, move it to a directory that's already in your system's PATH, such as /usr/local/bin
|
||||
|
||||
- ***Signup and get your authtoken and execute the following command:***
|
||||
```bash
|
||||
ngrok config add-authtoken $YOUR_AUTHTOKEN
|
||||
```
|
||||
|
||||
- Verify by running the **ngrok --version**
|
||||
|
||||
|
||||
4. **Expose Your Local Server with ngrok:**
|
||||
open another terminal and run ngrok to create a public URL that tunnels to your local port 8000.
|
||||
|
||||
```bash
|
||||
ngrok http 8000
|
||||
```
|
||||
ngrok will give you a public **Forwarding URL**, which will look something like **https://<random-string>.ngrok-free.dev**
|
||||
|
||||
|
||||
5. **Update Environment Variable:**
|
||||
You need to set the SERVICE_URL as an environment variable to the ngrok forwarding URL. Update your SERVICE_URL in your .env file
|
||||
This is critical because your application uses this URL to tell Twilio where to connect the WebSocket.
|
||||
```bash
|
||||
SERVICE_URL=https://<random-string>.ngrok-free.dev
|
||||
```
|
||||
|
||||
6. **Run the FastAPI App Locally:**
|
||||
Start the application on your another terminal in your local machine after updating .env file with SERVICE_URL.
|
||||
```bash
|
||||
uvicorn main:app --host 0.0.0.0 --port 8000
|
||||
```
|
||||
|
||||
7. **Set up a Twilio Trial Account:**
|
||||
- Create a free trial account at [Twilio](https://www.twilio.com/try-twilio).
|
||||
- Once your account is created, you will get a trial phone number and free credits to get started.
|
||||
|
||||
8. **Configure Twilio Webhook:**
|
||||
- Go to your Twilio phone number's configuration in the Twilio console.
|
||||
- Update the "A CALL COMES IN" webhook URL to point to your ngrok URL, followed by the /twiml endpoint (e.g., https://<random-string>.ngrok-free.dev/twiml).
|
||||
|
||||
Now, wait for around 2 mins and call your Twilio number, Twilio will send the webhook to the public ngrok URL, which will forward it to your locally running application. This allows you to test the entire end-to-end flow and see logs in real-time in your local terminal for debugging.
|
||||
|
||||
## Deployment on Google Cloud Platform
|
||||
|
||||
**Important Note on IAM Permissions:** Before deploying, ensure your Google Cloud account has the necessary IAM permissions for Cloud Run and Cloud Build Service account. Specifically, you will need roles such as `Cloud Run Invoker`,..etc
|
||||
|
||||
The `deploy.sh` script automates the process of building the container image, pushing it to Google Container Registry, and deploying it to Cloud Run.
|
||||
|
||||
|
||||
1. **Update `deploy.sh`:**
|
||||
- Open the `deploy.sh` file.
|
||||
- Replace `[YOUR_PROJECT_ID]` with your Google Cloud Project ID.
|
||||
|
||||
2. **Configure Docker (optional):**
|
||||
- Ensure you have Docker configured to work with `gcloud`:
|
||||
```bash
|
||||
gcloud auth configure-docker
|
||||
```
|
||||
|
||||
3. **Run the Deployment Script:**
|
||||
- Execute the `deploy.sh` script from your terminal:
|
||||
```bash
|
||||
bash deploy.sh
|
||||
```
|
||||
- The script will build and push the container, then deploy the service to Cloud Run. It will output the `Service URL` upon completion.
|
||||
|
||||
4. **Configure Twilio Webhook:**
|
||||
- Copy the `Service URL` from the output of the deployment script.
|
||||
- Go to your Twilio phone number's configuration in the Twilio Console.
|
||||
- Set the "A CALL COMES IN" webhook to the deployed service URL, followed by `/twiml` (e.g., `https://<your-service-url>.a.run.app/twiml`), and set the method to `HTTP POST`.
|
||||
- Save the Twilio configuration.
|
||||
|
||||
Your AI Care Assistant is now deployed and ready to receive calls.
|
||||
|
||||
Made by [Vishnu Vardhan Reddy Kanamata Reddy](https://github.com/KVishnuVardhanR)
|
||||
@@ -0,0 +1,90 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# https://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
|
||||
#!/bin/bash
|
||||
|
||||
# IMPORTANT: Before running this script, ensure you have:
|
||||
# 1. Authenticated with gcloud: `gcloud auth login`
|
||||
# 2. Configured Docker for Artifact Registry: `gcloud auth configure-docker`
|
||||
# 3. Replaced [YOUR_PROJECT_ID] and [YOUR_GEMINI_API_KEY] below.
|
||||
|
||||
# --- Step 1: Build and Push the Container Image ---
|
||||
# Replace [YOUR_PROJECT_ID] with your Google Cloud Project ID
|
||||
PROJECT_ID=[YOUR_PROJECT_ID]
|
||||
|
||||
# Generate a unique tag based on the current time (Forces a fresh pull)
|
||||
# --- Step 2: Deploy to Cloud Run ---
|
||||
|
||||
|
||||
SERVICE_NAME="gemini-live-health" # Cloud Run service name
|
||||
REGION="us-central1" # Cloud Run region
|
||||
|
||||
SERVICE_URL="" # Will be Updated with your deployed service URL after first deployment
|
||||
|
||||
TIMESTAMP=$(date +%Y%m%d-%H%M%S)
|
||||
IMAGE_NAME="gcr.io/${PROJECT_ID}/${SERVICE_NAME}:${TIMESTAMP}"
|
||||
|
||||
echo "Building and pushing container image: ${IMAGE_NAME}"
|
||||
gcloud builds submit --tag "${IMAGE_NAME}" --no-cache
|
||||
|
||||
if [ $? -ne 0 ]; then
|
||||
echo "Container build failed. Exiting."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
|
||||
|
||||
echo "Deploying service ${SERVICE_NAME} to Cloud Run in ${REGION}"
|
||||
DEPLOY_OUTPUT=$(gcloud run deploy "${SERVICE_NAME}" \
|
||||
--image "${IMAGE_NAME}" \
|
||||
--platform managed \
|
||||
--region "${REGION}" \
|
||||
--allow-unauthenticated \
|
||||
--min-instances=1 \
|
||||
--timeout=3600 \
|
||||
--memory=2Gi \
|
||||
--session-affinity \
|
||||
--concurrency=1 \
|
||||
--cpu=2 \
|
||||
--no-cpu-throttling \
|
||||
--set-env-vars="SERVICE_URL=${SERVICE_URL}" \
|
||||
--format="value(status.url)")
|
||||
|
||||
if [ $? -ne 0 ]; then
|
||||
echo "Cloud Run deployment failed. Exiting."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Updating service with its own public URL..."
|
||||
gcloud run services update "${SERVICE_NAME}" \
|
||||
--region "${REGION}" \
|
||||
--update-env-vars="SERVICE_URL=${DEPLOY_OUTPUT}"
|
||||
|
||||
if [ $? -ne 0 ]; then
|
||||
echo "Failed to update service with its public URL. Exiting."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
|
||||
# The SERVICE_URL is now directly set in the deployment command.
|
||||
# The DEPLOY_OUTPUT still contains the full URL, which is useful for Twilio configuration.
|
||||
|
||||
echo "Service deployed to: ${DEPLOY_OUTPUT}"
|
||||
echo "Deployment script completed successfully!"
|
||||
echo "Next steps:"
|
||||
echo "1. Copy the Service URL: ${DEPLOY_OUTPUT}"
|
||||
echo "2. Go to your Twilio phone number configuration in the Twilio Console."
|
||||
echo "3. Set the 'A CALL COMES IN' webhook to: ${DEPLOY_OUTPUT}/twiml (HTTP POST)"
|
||||
echo "4. Save the Twilio configuration and make a test call!"
|
||||
@@ -0,0 +1,191 @@
|
||||
# **Architecting a Low-Latency, High-Fidelity Conversational AI: A Definitive Guide to Integrating Twilio, FastAPI, and Gemini Live on Cloud Run**
|
||||
|
||||
## **Executive Summary**
|
||||
|
||||
This report provides a comprehensive architectural blueprint for a real-time, bidirectional voice-to-AI application. It addresses the three core technical challenges presented by the integration of Twilio, a FastAPI backend, and the Google Gemini Live API. These challenges are: (1) the system-level integration of Twilio's telephony with a custom backend via bidirectional WebSockets, (2) the design of a high-fidelity, low-latency audio transcoding pipeline to bridge the disparate audio formats (8kHz µ-law, 16kHz PCM, and 24kHz PCM), and (3) the deployment of this stateful, low-latency service on the stateless, serverless Google Cloud Run platform.
|
||||
|
||||
The analysis concludes that the most critical, non-obvious design decision is the selection of a *streaming* digital signal processing (DSP) library for audio resampling. Naive, chunk-by-chunk processing with standard libraries will introduce audible artifacts. The report identifies python-samplerate (a wrapper for libsamplerate) as the optimal choice, as its stateful "Full API" is specifically designed for high-quality, real-time chunked audio processing.
|
||||
|
||||
Furthermore, this document presents a production-ready Cloud Run architecture that addresses the platform's inherent challenges. It demonstrates that "excellent low latency" is achieved by mitigating cold starts using the min-instances=1 configuration. It also solves the critical state-management problem for a horizontally-scaling, stateless service by proposing a hybrid state model. This model utilizes Cloud Run's session affinity for in-memory DSP state.
|
||||
|
||||
**special note:** Having a Google Memorystore (Redis) for externalized conversational state, ensures a seamless, scalable, and low-latency conversational experience.
|
||||
|
||||
---
|
||||
|
||||
## **I. System Architecture: A High-Throughput Pipeline for Conversational AI**
|
||||
|
||||
This section details the high-level architecture of the complete system, outlining the components, their responsibilities, and the flow of data and control from the user's mobile phone to the Google Gemini Live API and back.
|
||||
|
||||
### **A. The End-to-End Data and Control Flow**
|
||||
|
||||
The entire system is orchestrated as a series of handoffs, transitioning from a standard HTTP webhook model to a persistent, bidirectional WebSocket stream. The lifecycle of a single interaction is as follows:
|
||||
|
||||
<img width="933" height="452" alt="image" src="https://github.com/user-attachments/assets/de89aa2b-c01f-4605-9cda-2ec8c32db994" />
|
||||
|
||||
|
||||
1. **Initiation (HTTP):** A user dials the Twilio phone number provisioned for the service.
|
||||
2. **Webhook Trigger:** Twilio's infrastructure receives the inbound call. Per its configuration, Twilio sends a synchronous HTTP POST request (a webhook) to the application's pre-defined HTTP endpoint (`/twiml`).
|
||||
3. **TwiML Response (HTTP):** The FastAPI server, deployed on Cloud Run, receives this HTTP request. It dynamically generates and returns a TwiML (Twilio Markup Language) document.
|
||||
4. **WebSocket Connection (WSS):** The TwiML response contains the `<Connect><Stream>` verb. Upon receiving this, Twilio's media servers initiate a persistent, secure WebSocket (WSS) connection to the application's WebSocket endpoint (`/ws/twilio`).
|
||||
5. **Bidirectional Streaming (WSS):** The WebSocket connection is established. The FastAPI server, using `asyncio`, manages two concurrent audio streams:
|
||||
* **Inbound Stream (User-to-AI):** Receiving 8kHz µ-law audio from Twilio, transcoding it in real-time to 16kHz PCM, and forwarding it to the Google Gemini Live API.
|
||||
* **Outbound Stream (AI-to-User):** Receiving 24kHz PCM audio from the Google Gemini Live API, transcoding it in real-time to 8kHz µ-law, and streaming it back to Twilio.
|
||||
6. **Termination:** The call concludes when the user hangs up or the connection is otherwise closed. The application cleans up resources, including the temporary transcription file.
|
||||
|
||||
### **B. Component 1: Twilio Programmable Voice & TwiML**
|
||||
|
||||
This component is the gateway between the public telephone network and the application.
|
||||
|
||||
#### **Phone Number Configuration**
|
||||
|
||||
The Twilio phone number is configured with a webhook that points to the `/twiml` endpoint of the deployed Cloud Run service.
|
||||
|
||||
#### **The `<Connect><Stream>` TwiML**
|
||||
|
||||
The application uses the `<Connect><Stream>` TwiML verb to establish a bidirectional WebSocket connection. The `/twiml` endpoint in `main.py` generates this TwiML dynamically, inserting the WebSocket URL of the service.
|
||||
|
||||
### **C. Component 2: FastAPI Backend on Cloud Run**
|
||||
|
||||
The core of the application is a FastAPI server deployed on Google Cloud Run.
|
||||
|
||||
#### **Endpoints**
|
||||
|
||||
* `/twiml` (POST): Receives the initial webhook from Twilio and responds with the TwiML to establish the WebSocket stream.
|
||||
* `/ws/twilio` (WebSocket): The main endpoint for the bidirectional audio stream. It orchestrates the flow of audio between Twilio and the Google Gemini Live API.
|
||||
|
||||
#### **Asynchronous Audio Handling**
|
||||
|
||||
The application leverages `asyncio` to handle the concurrent inbound and outbound audio streams. `asyncio.Queue` is used to pass audio chunks between the different processing tasks.
|
||||
|
||||
#### **Audio Transcoding Pipeline**
|
||||
|
||||
A critical part of the application is the real-time audio transcoding pipeline:
|
||||
|
||||
1. **Twilio to Gemini:**
|
||||
* The incoming base64-encoded µ-law audio from Twilio is decoded.
|
||||
* `audioop.ulaw2lin` converts the µ-law audio to 16-bit PCM.
|
||||
* The PCM audio is converted to a NumPy array of floats.
|
||||
* `samplerate.Resampler` upsamples the audio from 8kHz to 16kHz.
|
||||
* The resampled audio is converted back to 16-bit PCM bytes and sent to the Gemini Live API.
|
||||
2. **Gemini to Twilio:**
|
||||
* The 24kHz PCM audio from the Gemini Live API is received as bytes.
|
||||
* The bytes are converted to a NumPy array of floats.
|
||||
* `samplerate.Resampler` downsamples the audio from 24kHz to 8kHz.
|
||||
* The resampled audio is converted to 16-bit PCM.
|
||||
* `audioop.lin2ulaw` converts the PCM audio to µ-law.
|
||||
* The µ-law audio is base64-encoded and sent to Twilio.
|
||||
|
||||
### **D. Component 3: Google Gemini Live API**
|
||||
|
||||
The application uses the `google-genai` library to interact with the Google Gemini Live API via the Vertex AI platform.
|
||||
|
||||
#### **Session Management**
|
||||
|
||||
The `run_gemini_session` function in `utils/live_api.py` manages the persistent connection with the Gemini Live API. It:
|
||||
1. Utilizes `session_handle` for seamless session resumption, ensuring conversational context is maintained across potential connection drops.
|
||||
2. Constructs a `LiveConnectConfig` using `utils/live_api_config.py`, which includes system instructions, audio configuration, and real-time input settings (VAD).
|
||||
3. Establishes a single, persistent session for the duration of the call, with automatic reconnection logic if the stream is interrupted.
|
||||
4. Orchestrates two concurrent tasks (`sender_loop` and `heartbeat_loop`) and a primary message receiver to manage bidirectional audio and control signals.
|
||||
|
||||
### **E. Component 4: The `utils/` Package**
|
||||
|
||||
The logic is modularized into a `utils/` package containing specialized modules:
|
||||
|
||||
* **`live_api.py`**: Contains `run_gemini_session`, the core orchestrator for Gemini Live interactions, supporting sender/receiver loops and session resumption.
|
||||
* **`live_api_config.py`**: Defines the `LiveConnectConfig`, centralizing API settings, system instructions, and Voice Activity Detection (VAD) parameters.
|
||||
* **`audio_transcoding.py`**: Implements the real-time audio pipeline, handling conversions between Twilio's 8kHz µ-law and Gemini's 16kHz/24kHz PCM formats.
|
||||
* **`prompt.py`**: Stores the `BASE_SYSTEM_INSTRUCTION` and other persona-related configurations.
|
||||
|
||||
### **F. State Management**
|
||||
|
||||
For this sample application, state management is handled using an in-memory and handle-based approach:
|
||||
|
||||
* **In-memory state**: The `call_state` dictionary maintains the real-time status of the call, including activity flags and stream SIDs.
|
||||
* **Session Resumption Handle**: Instead of file-based history, the system leverages Gemini's `session_handle` property. This token is captured from session updates and used during reconnection to restore the full conversational context within the API itself.
|
||||
|
||||
For a production system, a more robust solution like **Google Memorystore (Redis)** would be recommended to externalize the conversational state, allowing for horizontal scaling and better resilience.
|
||||
|
||||
|
||||
# Implementation Plan:
|
||||
|
||||
A step-by-step process to create the Gemini Live Telephony application, a real-time, bidirectional voice-to-AI system integrating Twilio, FastAPI, and Google Gemini Live on Cloud Run.
|
||||
|
||||
## 1. Project Setup and Dependencies
|
||||
|
||||
- **Initialize Project:** Set up a Python project with a virtual environment.
|
||||
- **Install Dependencies:** Install necessary Python libraries as listed in `requirements.txt`:
|
||||
- `fastapi`: For the web server.
|
||||
- `uvicorn`: As the ASGI server.
|
||||
- `python-dotenv`: To manage environment variables.
|
||||
- `google-generativeai`: The Google Gemini SDK.
|
||||
- `numpy`: For numerical operations on audio data.
|
||||
- `samplerate`: For high-quality audio resampling.
|
||||
- `audioop`: For µ-law and linear PCM audio conversions.
|
||||
- **Environment Configuration:** Create a `.env` file to store information like `GOOGLE_CLOUD_PROJECT`, `SERVICE_URL` and other gemini configurations.
|
||||
|
||||
## 2. FastAPI Server Implementation (`main.py`)
|
||||
|
||||
- **Create `main.py`:** This will be the entry point for the FastAPI application.
|
||||
- **Implement `/twiml` Endpoint:**
|
||||
- Create an HTTP POST endpoint that responds to Twilio's webhook.
|
||||
- This endpoint will generate and return a TwiML response with the `<Connect><Stream>` verb, pointing to the WebSocket endpoint.
|
||||
- **Implement `/ws/twilio` WebSocket Endpoint:**
|
||||
- This will be the core of the application, handling the bidirectional audio stream.
|
||||
- It will manage the WebSocket lifecycle: connection, message handling, and disconnection.
|
||||
- It will use `asyncio.create_task` to run three concurrent tasks:
|
||||
1. `handle_twilio_to_gemini`: Processes audio from Twilio to Gemini.
|
||||
2. `handle_gemini_to_twilio`: Processes audio from Gemini to Twilio.
|
||||
3. `run_gemini_session`: Manages the high-level Gemini session logic from `utils/live_api.py`.
|
||||
|
||||
## 3. Audio Transcoding Pipeline (`utils/audio_transcoding.py`)
|
||||
|
||||
- **Inbound (Twilio to Gemini):**
|
||||
- **Base64 Decoding:** Decode the Base64 payload from Twilio media messages.
|
||||
- **µ-law to PCM:** Convert the 8kHz µ-law audio to 16-bit linear PCM using `audioop.ulaw2lin`.
|
||||
- **Upsampling:** Upsample the 8kHz PCM to 16kHz PCM using the `samplerate` library.
|
||||
- **Outbound (Gemini to Twilio):**
|
||||
- **Downsampling:** Downsample the 24kHz PCM audio from Gemini to 8kHz PCM using `samplerate`.
|
||||
- **PCM to µ-law:** Convert the 8kHz linear PCM to µ-law using `audioop.lin2ulaw`.
|
||||
- **Base64 Encoding:** Encode the µ-law audio to Base64 to be sent back to Twilio.
|
||||
|
||||
## 4. Gemini Live Integration (`utils/live_api.py`)
|
||||
|
||||
- **Initialize Gemini Client**: Set up the Gemini client with Vertex AI and project/location settings.
|
||||
- **`run_gemini_session` Function**:
|
||||
- **`LiveConnectConfig`**: Configure the session via `utils/live_api_config.py`.
|
||||
- **Session Management**: Establish a persistent connection that stays active throughout the call.
|
||||
- **Sender/Heartbeat/Receiver Loops**: Use concurrent tasks to stream audio and monitor for session updates.
|
||||
- **Resumption Logic**: Utilize `session_handle` to restore state during reconnection.
|
||||
|
||||
## 5. Utility Package (`utils/`)
|
||||
|
||||
- **Modular Organization**: The code is separated into multiple files within the `utils/` directory.
|
||||
- **`live_api_config.py`**: Manages all connection parameters, including VAD sensitivity and system instructions.
|
||||
- **`prompt.py`**: Centralizes the `BASE_SYSTEM_INSTRUCTION` for the AI persona.
|
||||
- **`audio_transcoding.py`**: Contains the resampling and transcoding logic previously in `main.py`.
|
||||
|
||||
## 6. State Management
|
||||
|
||||
- **In-memory State**: Use a Python dictionary (`call_state`) to manage the active state of the call within a single Cloud Run instance.
|
||||
- **Session Resumption Logic**:
|
||||
- Instead of file-based history, captures `session_resumption_update` messages.
|
||||
- Saves the `new_handle` to be reused in subsequent `connect` calls.
|
||||
- **Note on Production Systems**: For a production-grade application, externalize state to a robust solution like Google Memorystore (Redis).
|
||||
|
||||
## 7. Containerization and Deployment
|
||||
|
||||
- **Create `Dockerfile`:**
|
||||
- Use a slim Python base image.
|
||||
- Install system dependencies like `libsamplerate0`.
|
||||
- Install Python dependencies from `requirements.txt`.
|
||||
- Use `uvicorn` to run the application.
|
||||
- **Create `deploy.sh`:**
|
||||
- Write a shell script to automate the process of building the Docker image and deploying it to Google Cloud Run.
|
||||
- **Deploy to Cloud Run:**
|
||||
- Configure the Cloud Run service with the following settings, as seen in `deploy.sh`:
|
||||
- `--min-instances=1`: This is crucial for a low-latency application to avoid "cold starts," ensuring that an instance is always running and ready to accept calls.
|
||||
- `--timeout=3600`: Sets a long request timeout (1 hour) to accommodate long-running WebSocket connections for phone calls.
|
||||
- `--memory=2Gi` and `--cpu=2`: Allocates sufficient resources for the CPU-intensive audio resampling tasks.
|
||||
- `--session-affinity`: Ensures that requests from the same client (Twilio, in this case) are routed to the same Cloud Run instance, which is important for maintaining the WebSocket connection.
|
||||
- `--concurrency=1`: This is a critical setting for this application. Since a single instance handles a single, stateful phone call at a time, setting concurrency to 1 ensures that each instance is dedicated to a single call. This prevents issues with managing multiple concurrent calls on a single instance.
|
||||
- `--no-cpu-throttling`: The audio resampling process is CPU-intensive and sensitive to latency. Disabling CPU throttling ensures that the instance has access to the full allocated CPU, which is essential for real-time audio processing and maintaining a smooth, low-latency conversation.
|
||||
@@ -0,0 +1,94 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# https://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
|
||||
import samplerate # pip install samplerate
|
||||
import uvicorn
|
||||
from dotenv import load_dotenv
|
||||
from fastapi import FastAPI, Response, WebSocket
|
||||
from google import genai
|
||||
from utils.audio_transcoding import handle_gemini_to_twilio, handle_twilio_to_gemini
|
||||
from utils.live_api import run_gemini_session
|
||||
|
||||
load_dotenv()
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
logging.getLogger("websockets").setLevel(logging.WARNING)
|
||||
logging.getLogger("google.auth").setLevel(logging.WARNING)
|
||||
|
||||
app = FastAPI(title="Gemini Live Health Demo")
|
||||
|
||||
# --- CONFIGURATION ---
|
||||
MODEL_ID = os.getenv("GOOGLE_GENAI_MODEL", "gemini-live-2.5-flash-native-audio")
|
||||
|
||||
# Initialize Gemini Client
|
||||
try:
|
||||
client = genai.Client(
|
||||
vertexai=True,
|
||||
project=os.getenv("GOOGLE_CLOUD_PROJECT"),
|
||||
location=os.getenv("GOOGLE_CLOUD_LOCATION"),
|
||||
)
|
||||
# client = genai.Client(api_key=os.environ["GEMINI_API_KEY"])
|
||||
except KeyError:
|
||||
logger.fatal("Google Cloud configuration not found in environment variables.")
|
||||
exit(1)
|
||||
|
||||
|
||||
@app.post("/twiml")
|
||||
async def get_twiml():
|
||||
"""Generates TwiML response to initiate a WebSocket stream with Twilio."""
|
||||
service_url = (
|
||||
os.getenv("SERVICE_URL").replace("https://", "").replace("http://", "")
|
||||
)
|
||||
twiml = f"""<Response><Connect><Stream url="wss://{service_url}/ws/twilio" /></Connect></Response>"""
|
||||
return Response(content=twiml, media_type="application/xml")
|
||||
|
||||
|
||||
@app.websocket("/ws/twilio")
|
||||
async def websocket_twilio_endpoint(websocket: WebSocket):
|
||||
"""The main WebSocket endpoint for handling Twilio media streams."""
|
||||
await websocket.accept()
|
||||
call_state = {"active": False}
|
||||
in_q, out_q = asyncio.Queue(), asyncio.Queue()
|
||||
resampler_in = samplerate.Resampler("sinc_fastest", channels=1)
|
||||
resampler_out = samplerate.Resampler("sinc_fastest", channels=1)
|
||||
|
||||
tasks = [
|
||||
asyncio.create_task(
|
||||
handle_twilio_to_gemini(websocket, in_q, resampler_in, call_state)
|
||||
),
|
||||
asyncio.create_task(
|
||||
handle_gemini_to_twilio(websocket, out_q, resampler_out, call_state)
|
||||
),
|
||||
asyncio.create_task(
|
||||
run_gemini_session(client, MODEL_ID, in_q, out_q, call_state)
|
||||
),
|
||||
]
|
||||
|
||||
try:
|
||||
await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
|
||||
finally:
|
||||
call_state["active"] = False
|
||||
for t in tasks:
|
||||
t.cancel()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
uvicorn.run(app, host="0.0.0.0", port=8000)
|
||||
@@ -0,0 +1,9 @@
|
||||
fastapi==0.111.0
|
||||
uvicorn[standard]==0.30.1
|
||||
gunicorn==23.0.0
|
||||
python-dotenv==1.0.1
|
||||
google-genai==1.28.0
|
||||
twilio==9.8.6
|
||||
websockets==15.0.1
|
||||
numpy==2.3.4
|
||||
samplerate==0.2.2
|
||||
@@ -0,0 +1,112 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# https://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import asyncio
|
||||
import audioop
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
|
||||
import numpy as np
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def handle_twilio_to_gemini(
|
||||
websocket, audio_queue: asyncio.Queue, resampler, call_state
|
||||
):
|
||||
"""Handles the inbound audio stream from Twilio to the Gemini API."""
|
||||
async for message_str in websocket.iter_text():
|
||||
try:
|
||||
await asyncio.sleep(0) # Yield control to event loop
|
||||
msg = json.loads(message_str)
|
||||
if msg["event"] == "media":
|
||||
if not call_state.get("active"):
|
||||
continue
|
||||
|
||||
# 1. Decode Base64
|
||||
chunk_ulaw = base64.b64decode(msg["media"]["payload"])
|
||||
|
||||
# 2. Decode u-law -> PCM
|
||||
chunk_pcm = audioop.ulaw2lin(chunk_ulaw, 2)
|
||||
|
||||
# 3. PCM -> Float32
|
||||
arr_8k = np.frombuffer(chunk_pcm, dtype=np.int16)
|
||||
arr_8k_float = arr_8k.astype(np.float32) / 32768.0
|
||||
|
||||
# 4. Resample 8k -> 16k
|
||||
arr_16k_float = resampler.process(
|
||||
arr_8k_float, ratio=2.0, end_of_input=False
|
||||
)
|
||||
|
||||
# 5. Float32 -> Int16
|
||||
arr_16k = (arr_16k_float * 32767).astype(np.int16)
|
||||
|
||||
await audio_queue.put(arr_16k.tobytes())
|
||||
|
||||
elif msg["event"] == "start":
|
||||
call_state["stream_sid"] = msg["start"]["streamSid"]
|
||||
call_state["active"] = True
|
||||
logger.info(f"Stream started: {msg['start']['streamSid']}")
|
||||
elif msg["event"] == "stop":
|
||||
call_state["active"] = False
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Inbound error: {e}")
|
||||
break
|
||||
|
||||
|
||||
async def handle_gemini_to_twilio(
|
||||
websocket, audio_queue: asyncio.Queue, resampler, call_state
|
||||
):
|
||||
"""Handles the outbound audio stream from the Gemini API to Twilio."""
|
||||
logger.info("--- RUNNING LATEST VERSION OF gemini_to_twilio ---")
|
||||
while True:
|
||||
await asyncio.sleep(0)
|
||||
try:
|
||||
# Receive 24k PCM bytes from Gemini
|
||||
chunk_24k_bytes = await asyncio.wait_for(audio_queue.get(), timeout=1.0)
|
||||
|
||||
if chunk_24k_bytes:
|
||||
# 1. Bytes -> Float32
|
||||
arr_24k = np.frombuffer(chunk_24k_bytes, dtype=np.int16)
|
||||
arr_24k_float = arr_24k.astype(np.float32) / 32768.0
|
||||
|
||||
# 2. Resample 24k -> 8k
|
||||
arr_8k_float = resampler.process(
|
||||
arr_24k_float, ratio=(8000 / 24000), end_of_input=False
|
||||
)
|
||||
|
||||
# 3. Float32 -> Int16
|
||||
arr_8k = (arr_8k_float * 32767).astype(np.int16)
|
||||
|
||||
# 4. PCM -> u-law
|
||||
chunk_ulaw = audioop.lin2ulaw(arr_8k.tobytes(), 2)
|
||||
|
||||
# 5. Send
|
||||
payload = base64.b64encode(chunk_ulaw).decode("utf-8")
|
||||
if sid := call_state.get("stream_sid"):
|
||||
await websocket.send_json(
|
||||
{
|
||||
"event": "media",
|
||||
"streamSid": sid,
|
||||
"media": {"payload": payload},
|
||||
}
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
if not call_state.get("active", True):
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Outbound error: {e}")
|
||||
continue
|
||||
@@ -0,0 +1,142 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# https://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
|
||||
from google.genai import types
|
||||
from utils.prompt import BASE_SYSTEM_INSTRUCTION
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def run_gemini_session(client, model_id, in_q, out_q, call_state):
|
||||
"""Handles Gemini flow with persistent connection and session extension."""
|
||||
while not call_state.get("active"):
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
session_handle = None
|
||||
|
||||
# 10ms of silence at 16kHz (320 bytes for 16-bit PCM)
|
||||
silent_chunk = b"\x00" * 320
|
||||
|
||||
while call_state.get("active"):
|
||||
try:
|
||||
logger.info(
|
||||
f"Connecting to Gemini (Resumption: {session_handle is not None}, Handle: {session_handle})..."
|
||||
)
|
||||
|
||||
config = types.LiveConnectConfig(
|
||||
system_instruction=types.Content(
|
||||
parts=[types.Part(text=BASE_SYSTEM_INSTRUCTION)]
|
||||
),
|
||||
response_modalities=["AUDIO"],
|
||||
session_resumption=types.SessionResumptionConfig(handle=session_handle),
|
||||
speech_config=types.SpeechConfig(
|
||||
voice_config=types.VoiceConfig(
|
||||
prebuilt_voice_config=types.PrebuiltVoiceConfig(
|
||||
voice_name="Achird",
|
||||
)
|
||||
),
|
||||
language_code="en-US",
|
||||
),
|
||||
realtime_input_config=types.RealtimeInputConfig(
|
||||
automatic_activity_detection=types.AutomaticActivityDetection(
|
||||
disabled=False,
|
||||
start_of_speech_sensitivity=types.StartSensitivity.START_SENSITIVITY_LOW,
|
||||
end_of_speech_sensitivity=types.EndSensitivity.END_SENSITIVITY_LOW,
|
||||
prefix_padding_ms=20,
|
||||
silence_duration_ms=150,
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
async with client.aio.live.connect(
|
||||
model=model_id, config=config
|
||||
) as session:
|
||||
logger.info("Gemini Connected.")
|
||||
|
||||
async def sender_loop():
|
||||
while call_state.get("active"):
|
||||
try:
|
||||
chunk = await asyncio.wait_for(in_q.get(), timeout=0.01)
|
||||
await session.send_realtime_input(
|
||||
audio=types.Blob(
|
||||
data=chunk, mime_type="audio/pcm;rate=16000"
|
||||
)
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
except Exception as e:
|
||||
logger.error(f"Sender error: {e}")
|
||||
break
|
||||
|
||||
async def heartbeat_loop():
|
||||
while call_state.get("active"):
|
||||
try:
|
||||
await session.send_realtime_input(
|
||||
audio=types.Blob(
|
||||
data=silent_chunk, mime_type="audio/pcm;rate=16000"
|
||||
)
|
||||
)
|
||||
await asyncio.sleep(5)
|
||||
except Exception as e:
|
||||
logger.error(f"Heartbeat error: {e}")
|
||||
break
|
||||
|
||||
sender_task = asyncio.create_task(sender_loop())
|
||||
heartbeat_task = asyncio.create_task(heartbeat_loop())
|
||||
|
||||
while call_state.get("active"):
|
||||
try:
|
||||
message = await asyncio.wait_for(
|
||||
session.receive().__anext__(), timeout=0.01
|
||||
)
|
||||
|
||||
if message.session_resumption_update:
|
||||
update = message.session_resumption_update
|
||||
if update.new_handle:
|
||||
session_handle = update.new_handle
|
||||
logger.info(f"!!! SAVED HANDLE: {session_handle} !!!")
|
||||
|
||||
if message.server_content:
|
||||
if message.server_content.model_turn:
|
||||
for part in message.server_content.model_turn.parts:
|
||||
if part.inline_data:
|
||||
await out_q.put(part.inline_data.data)
|
||||
|
||||
if message.server_content.turn_complete:
|
||||
logger.info(
|
||||
"[Session] Turn complete. Keeping session alive..."
|
||||
)
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
except StopAsyncIteration:
|
||||
logger.warning("[Session] Server closed the stream.")
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Receiver error: {e}")
|
||||
break
|
||||
|
||||
sender_task.cancel()
|
||||
heartbeat_task.cancel()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Session error: {e}")
|
||||
if not call_state.get("active"):
|
||||
break
|
||||
await asyncio.sleep(2)
|
||||
|
||||
logger.info("Session cycle completed.")
|
||||
@@ -0,0 +1,30 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# https://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
BASE_SYSTEM_INSTRUCTION = """
|
||||
You are Sam, an AI care team member representing Northwestern Medicine, reaching out to a patient named Vishnu Vardhan via a phone call.
|
||||
|
||||
Context:
|
||||
The patient, Vishnu, recently completed his annual checkup. Your specific goal is to follow up on that visit to see how he is doing and ask if he would like to schedule any further visits or specialist follow-ups based on that appointment.
|
||||
|
||||
Instructions:
|
||||
1. **Persona:** Maintain a helpful, informative, and respectful tone. Your voice should be human-like, empathetic, and professional.
|
||||
2. **Interaction Style:** Build a natural, turn-taking dialogue. Listen carefully to Vishnu's responses and adapt your replies accordingly.
|
||||
3. **Objective:** Empower the patient to take charge of his health. Find out if he has outstanding questions or needs help booking next steps.
|
||||
4. **Handling Declines/Positive Health Status:** If Vishnu indicates that he feels fine or does not wish to schedule any further visits, you must accept this answer without pressure. Respond with "Good to know," say "Thank you," and politely end the call.
|
||||
5. **Constraints:** Strictly adhere to all WON'T constraints (e.g., do not provide medical diagnoses, do not be pushy, do not hallucinate appointments).
|
||||
|
||||
Opening Line:
|
||||
"Hi, this is Sam calling from the care team at Northwestern Medicine. Am I speaking with Vishnu Vardhan?"
|
||||
"""
|
||||
@@ -0,0 +1,3 @@
|
||||
[flake8]
|
||||
ignore = W291, W293, E261
|
||||
max-line-length = 100
|
||||
@@ -0,0 +1,8 @@
|
||||
[MESSAGES CONTROL]
|
||||
disable=
|
||||
E3701,
|
||||
W0613,
|
||||
C0302,
|
||||
C0301,
|
||||
R0903,
|
||||
too-many-lines
|
||||
@@ -0,0 +1 @@
|
||||
web: gunicorn --bind :8080 main:me
|
||||
@@ -0,0 +1,107 @@
|
||||
# Mesop application using Gemini API in Vertex AI on Cloud Run
|
||||
|
||||
| | |
|
||||
| --------- | --------------------------------------------- |
|
||||
| Author(s) | [Hussain Chinoy](https://github.com/ghchinoy) |
|
||||
|
||||
<!-- markdownlint-disable MD036 -->
|
||||
|
||||
**YouTube Video: How to build a Gemini powered Mesop app**
|
||||
|
||||
<!-- markdownlint-enable MD036 -->
|
||||
|
||||
<!-- markdownlint-disable MD033 -->
|
||||
<a href="https://www.youtube.com/watch?v=KUfPiSUJrwE&list=PLIivdWyY5sqJio2yeg1dlfILOUO2FoFRx" target="_blank">
|
||||
<img src="https://img.youtube.com/vi/KUfPiSUJrwE/maxresdefault.jpg" alt="How to build a Gemini powered Mesop app" width="500">
|
||||
</a>
|
||||
<!-- markdownlint-enable MD033 -->
|
||||
|
||||
This application demonstrates a [Mesop](https://github.com/google/mesop) UI framework application running on Cloud Run.
|
||||
|
||||
Sample screenshots and video demos of the application are shown below:
|
||||
|
||||
## Application screenshots
|
||||
|
||||

|
||||
|
||||
## Run the Application locally (on Cloud Shell)
|
||||
|
||||
> NOTE: **Before you move forward, ensure that you have followed the instructions in [SETUP.md](../SETUP.md).**
|
||||
> Additionally, ensure that you have cloned this repository and you are currently in the `gemini-mesop-cloudrun` folder. This should be your active working directory for the rest of the commands.
|
||||
|
||||
To run the Mesop application locally (on Cloud Shell), we need to perform the following steps:
|
||||
|
||||
1. Set up the Python virtual environment and install the dependencies:
|
||||
|
||||
In Cloud Shell, execute the following commands:
|
||||
|
||||
```bash
|
||||
python3 -m venv gemini-mesop
|
||||
. gemini-mesop/bin/activate
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
2. Your application requires access to two environment variables:
|
||||
|
||||
- `GOOGLE_CLOUD_PROJECT` : This the Google Cloud project ID.
|
||||
- `GOOGLE_CLOUD_REGION` : This is the region in which you are deploying your Cloud Run app. For e.g. us-central1.
|
||||
|
||||
These variables are needed since the Vertex AI initialization needs the Google Cloud project ID and the region.
|
||||
|
||||
In Cloud Shell, execute the following commands:
|
||||
|
||||
```bash
|
||||
export GOOGLE_CLOUD_PROJECT=$(gcloud config get project) # this will populate your current project ID
|
||||
export GOOGLE_CLOUD_REGION='us-central1' # If you change this, make sure the region is supported.
|
||||
```
|
||||
|
||||
3. To run the application locally, execute the following command:
|
||||
|
||||
In Cloud Shell, execute the following command:
|
||||
|
||||
```bash
|
||||
mesop --port 8080 main.py
|
||||
```
|
||||
|
||||
The application will start up and you will be provided a URL to the application. Use Cloud Shell's [web preview](https://cloud.google.com/shell/docs/using-web-preview) function to launch the preview page. You may also visit that in the browser to view the application. Choose the functionality that you would like to check out and the application will prompt the Gemini API in Vertex AI and display the responses.
|
||||
|
||||
## Build and Deploy the Application to Cloud Run
|
||||
|
||||
To deploy the Mesop Application in [Cloud Run](https://cloud.google.com/run/docs/quickstarts/deploy-container), we need to perform the following steps:
|
||||
|
||||
1. Your Cloud Run app requires access to two environment variables:
|
||||
|
||||
- `GOOGLE_CLOUD_PROJECT` : This the Google Cloud project ID.
|
||||
- `GOOGLE_CLOUD_REGION` : This is the region in which you are deploying your Cloud Run app. For e.g. us-central1.
|
||||
|
||||
These variables are needed since the Vertex AI initialization needs the Google Cloud project ID and the region.
|
||||
|
||||
In Cloud Shell, execute the following commands:
|
||||
|
||||
```bash
|
||||
export GOOGLE_CLOUD_PROJECT=$(gcloud config get project) # Use this or manually change this
|
||||
export GOOGLE_CLOUD_REGION='us-central1' # If you change this, make sure the region is supported.
|
||||
```
|
||||
|
||||
2. Build and deploy the service to Cloud Run:
|
||||
|
||||
In Cloud Shell, execute the following command to name the Cloud Run service:
|
||||
|
||||
```bash
|
||||
export SERVICE_NAME='mesop-gemini' # this is the name of our Application and Cloud Run service. Change this if you'd like to.
|
||||
```
|
||||
|
||||
In Cloud Shell, execute the following command:
|
||||
|
||||
```bash
|
||||
gcloud run deploy $SERVICE_NAME \
|
||||
--source . \
|
||||
--port=8080 --allow-unauthenticated \
|
||||
--project=$GOOGLE_CLOUD_PROJECT --region=$GOOGLE_CLOUD_REGION \
|
||||
--set-env-vars=GOOGLE_CLOUD_PROJECT=$GOOGLE_CLOUD_PROJECT \
|
||||
--set-env-vars=GOOGLE_CLOUD_REGION=$GOOGLE_CLOUD_REGION
|
||||
```
|
||||
|
||||
On successful deployment, you will be provided a URL to the Cloud Run service. You can visit that in the browser to view the Cloud Run application that you just deployed. Choose the functionality that you would like to check out and the application will prompt the Gemini API in Vertex AI and display the responses.
|
||||
|
||||
Congratulations!
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,9 @@
|
||||
[mypy]
|
||||
|
||||
check_untyped_defs = True
|
||||
disallow_untyped_defs = True
|
||||
|
||||
# stricter TypedDict handling
|
||||
disallow_any_generics = True
|
||||
warn_no_return = True
|
||||
no_implicit_reexport = True
|
||||
@@ -0,0 +1,4 @@
|
||||
mesop
|
||||
gunicorn==23.0.0
|
||||
dataclasses_json
|
||||
google-genai
|
||||
@@ -0,0 +1,3 @@
|
||||
[flake8]
|
||||
ignore = W291, W293
|
||||
max-line-length = 100
|
||||
@@ -0,0 +1,9 @@
|
||||
[mypy]
|
||||
|
||||
check_untyped_defs = True
|
||||
disallow_untyped_defs = True
|
||||
|
||||
# stricter TypedDict handling
|
||||
disallow_any_generics = True
|
||||
warn_no_return = True
|
||||
no_implicit_reexport = True
|
||||
@@ -0,0 +1,90 @@
|
||||
# Copyright 2024 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from dataclasses import field
|
||||
from typing import Any, Generator, TypedDict
|
||||
|
||||
from dataclasses_json import dataclass_json
|
||||
import mesop as me
|
||||
|
||||
# pylint: disable=E0402
|
||||
from .styles import _STYLE_CURRENT_NAV, _STYLE_MAIN_HEADER
|
||||
|
||||
# pylint: disable=E0402
|
||||
|
||||
|
||||
class Page(TypedDict):
|
||||
"""Page class"""
|
||||
|
||||
display: str
|
||||
route: str
|
||||
|
||||
|
||||
page_data = [
|
||||
{"display": "Generate story", "route": "/"},
|
||||
{"display": "Marketing campaign", "route": "/marketing"},
|
||||
{"display": "Image playground", "route": "/images"},
|
||||
{"display": "Video playground", "route": "/videos"},
|
||||
]
|
||||
|
||||
page_json = [Page(**data) for data in page_data] # type: ignore[typeddict-item]
|
||||
|
||||
|
||||
@dataclass_json
|
||||
@me.stateclass
|
||||
class State:
|
||||
"""Mesop state class"""
|
||||
|
||||
# pylint: disable=E3701
|
||||
pages: list[Page] = field(default_factory=lambda: page_json)
|
||||
current_page: str = ""
|
||||
# pylint: disable=E3701
|
||||
|
||||
|
||||
def navigate_to(e: me.ClickEvent) -> Generator[None, Any, None]:
|
||||
"""Navigate to a page event"""
|
||||
s = me.state(State)
|
||||
s.current_page = e.key
|
||||
me.navigate(e.key)
|
||||
yield
|
||||
|
||||
|
||||
def page_navigation_menu(url: str) -> None:
|
||||
"""Page navigation menu creation"""
|
||||
print(f"url: {url}")
|
||||
state = me.state(State)
|
||||
with me.box(style=_STYLE_MAIN_HEADER):
|
||||
with me.box(style=me.Style(display="flex", flex_direction="row", gap=12)):
|
||||
for page in state.pages:
|
||||
disabled = False
|
||||
if state.current_page == page.get("route"):
|
||||
disabled = True
|
||||
me.button(
|
||||
page.get("display"),
|
||||
key=f"{page.get('route')}",
|
||||
on_click=navigate_to,
|
||||
disabled=disabled,
|
||||
style=_STYLE_CURRENT_NAV if disabled else me.Style(),
|
||||
# type="flat" if disabled else "stroked"
|
||||
)
|
||||
|
||||
|
||||
@me.content_component
|
||||
def nav_menu(url: str) -> str:
|
||||
"""Navigation menu component"""
|
||||
page_navigation_menu(url=url)
|
||||
me.slot()
|
||||
|
||||
state = me.state(State)
|
||||
return state.current_page
|
||||
@@ -0,0 +1,39 @@
|
||||
# Copyright 2024 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""
|
||||
This is a set of prompts used in the Mesop app
|
||||
"""
|
||||
|
||||
# video
|
||||
|
||||
VIDEO_TAGS_PROMPT = """Answer the following questions using the video only:
|
||||
1. What is in the video?
|
||||
2. What objects are in the video?
|
||||
3. What is the action in the video?
|
||||
4. Provide 5 best tags for this video?
|
||||
|
||||
Give the answer in the table format with question and answer as columns.
|
||||
""" # noqa: E261, W291
|
||||
|
||||
|
||||
VIDEO_GEOLOCATION_PROMPT = """Answer the following questions using the video only:
|
||||
|
||||
What is this video about?
|
||||
How do you know which city it is?
|
||||
What street is this?
|
||||
What is the nearest intersection?
|
||||
|
||||
Answer the questions in a table format with question and answer as columns.
|
||||
""" # noqa: E261, W291
|
||||
@@ -0,0 +1,117 @@
|
||||
# Copyright 2024 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import mesop as me
|
||||
|
||||
_DEFAULT_BORDER = me.Border.all(
|
||||
me.BorderSide(
|
||||
color="#e0e0e0",
|
||||
width=1,
|
||||
style="solid",
|
||||
)
|
||||
)
|
||||
|
||||
_STYLE_MAIN_HEADER = me.Style(
|
||||
border=_DEFAULT_BORDER,
|
||||
padding=me.Padding.all(5),
|
||||
)
|
||||
|
||||
_STYLE_CURRENT_NAV = me.Style(color="#99000", border_radius=0, font_weight="bold")
|
||||
|
||||
|
||||
_DEFAULT_BORDER = me.Border.all(
|
||||
me.BorderSide(
|
||||
color="#e0e0e0",
|
||||
width=1,
|
||||
style="solid",
|
||||
)
|
||||
)
|
||||
|
||||
_STYLE_CONTAINER = me.Style(
|
||||
display="grid",
|
||||
grid_template_columns="5fr 2fr",
|
||||
grid_template_rows="auto 5fr",
|
||||
height="100vh",
|
||||
)
|
||||
|
||||
_STYLE_MAIN_HEADER = me.Style(
|
||||
border=_DEFAULT_BORDER, padding=me.Padding(top=15, left=15, right=15, bottom=5)
|
||||
)
|
||||
|
||||
_STYLE_MAIN_COLUMN = me.Style(
|
||||
border=_DEFAULT_BORDER,
|
||||
padding=me.Padding.all(15),
|
||||
overflow_y="scroll",
|
||||
)
|
||||
|
||||
_STYLE_TITLE_BOX = me.Style(display="inline-block")
|
||||
|
||||
_STORY_INPUT_STYLE = me.Style(
|
||||
width="500px"
|
||||
# display="flex",
|
||||
# flex_basis="max(100vh, calc(50% - 48px))",
|
||||
)
|
||||
|
||||
_BOX_STYLE = me.Style(
|
||||
flex_basis="max(100vh, calc(50% - 48px))",
|
||||
background="#fff",
|
||||
border_radius=12,
|
||||
box_shadow=("0 3px 1px -2px #0003, 0 2px 2px #00000024, 0 1px 5px #0000001f"),
|
||||
padding=me.Padding(top=16, left=16, right=16, bottom=16),
|
||||
display="flex",
|
||||
flex_direction="column",
|
||||
)
|
||||
|
||||
_SPINNER_STYLE = me.Style(
|
||||
display="flex",
|
||||
flex_direction="row",
|
||||
padding=me.Padding.all(16),
|
||||
align_items="center",
|
||||
gap=10,
|
||||
)
|
||||
|
||||
FANCY_TEXT_GRADIENT = me.Style(
|
||||
color="transparent",
|
||||
background=(
|
||||
"linear-gradient(72.83deg,#4285f4 11.63%,#9b72cb 40.43%,#d96570 68.07%)" " text"
|
||||
),
|
||||
)
|
||||
|
||||
_STYLE_CURRENT_TAB = me.Style(
|
||||
color="#99000",
|
||||
border_radius=0,
|
||||
font_weight="bold",
|
||||
border=me.Border(
|
||||
bottom=me.BorderSide(color="#000", width=2, style="solid"),
|
||||
top=None,
|
||||
right=None,
|
||||
left=None,
|
||||
),
|
||||
)
|
||||
|
||||
_STYLE_OTHER_TAB = me.Style(
|
||||
color="#8d8e9d",
|
||||
border_radius=0,
|
||||
# font_weight="bold",
|
||||
)
|
||||
|
||||
_TABBER_STYLE = me.Style(
|
||||
padding=me.Padding(top=0, right=0, left=0, bottom=2),
|
||||
border=me.Border(
|
||||
bottom=me.BorderSide(color="#e5e5e5", width=1, style="solid"),
|
||||
top=None,
|
||||
right=None,
|
||||
left=None,
|
||||
),
|
||||
)
|
||||
@@ -0,0 +1,235 @@
|
||||
# Non-blocking Chat app with Quart + Gemini Live API + Cloud Run
|
||||
|
||||
| | |
|
||||
| --------- | ------------------------------------------ |
|
||||
| Author(s) | [Kaz Sato](https://github.com/kazunori279) |
|
||||
|
||||
This application demonstrates a non-blocking communication with [Quart](https://quart.palletsprojects.com/en/latest/) and Gemini Live API running on Cloud Run.
|
||||
|
||||
## Application screenshot
|
||||
|
||||

|
||||
|
||||
Interruption example with the demo chat app
|
||||
|
||||
## Design Concepts
|
||||
|
||||
### Why Quart + Gemini Live API?
|
||||
|
||||
[Quart](https://quart.palletsprojects.com/en/latest/) is an asynchronous Python web framework built upon the ASGI standard, designed to facilitate the development of high-performance, concurrent applications. Its architecture and feature set render it particularly well-suited for constructing sophisticated generative AI applications that leverage real-time communication technologies like WebSockets and [Gemini Live API](https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/multimodal-live).
|
||||
|
||||
**Key Benefits of Quart:**
|
||||
|
||||
- **Asynchronous Architecture:** Quart's foundation in `asyncio` enables efficient handling of concurrent I/O-bound operations, crucial for interacting with AI models and managing real-time data streams without performance degradation.
|
||||
- **Native WebSocket Support:** The framework offers robust, integrated support for WebSockets, enabling persistent, bidirectional communication channels essential for interactive AI applications requiring real-time data exchange.
|
||||
- **Flask-Inspired API:** Quart's API design, mirroring the widely adopted Flask framework, promotes rapid development and leverages a familiar paradigm for developers, reducing the learning curve.
|
||||
- **Optimized for Multimodal Streaming Data:** The framework is engineered to process and transmit large data streams efficiently, a vital capability when dealing with the potentially voluminous multimodal outputs of generative AI models.
|
||||
|
||||
**Key Benefits building Gen AI app with Quart + Gemini Live API:**
|
||||
|
||||
- **Responsiveness and Natural Conversation:** Quart supports non-blocking, full-duplex WebSocket communication natively, crucial for a truly interactive Gen AI experience. It doesn't halt while waiting for Gemini, ensuring quick replies and a smooth conversation flow, especially when the app supports multimodal interaction using audio and images and is network-latency sensitive. Users can send text or voice messages in quick succession, and Quart handles them and interrupts with less delays.
|
||||
- **Concurrency and Scalability:** Handles many users and their messages simultaneously. Quart can process multiple requests and replies with Gemini concurrently, making the gen AI app faster and more efficient. Quart makes better use of server resources with the single thread event-loop design, leading to lower operational costs and better scalability.
|
||||
|
||||
### Flask (blocking) v. Quart (non-blocking)
|
||||
|
||||
**How Flask works:**
|
||||
|
||||
- Blocking: Flask handles one request at a time. It blocks while waiting for Gemini, causing delays. The diagram shows Flask "blocked" while waiting for a response.
|
||||
- Sequential: The client must wait for each response before sending the next message, making the interaction slow.
|
||||
|
||||

|
||||
|
||||
<!-- mermaid code:
|
||||
sequenceDiagram
|
||||
participant Client
|
||||
participant Flask
|
||||
participant Gemini
|
||||
|
||||
Client->>Flask: hello
|
||||
activate Flask
|
||||
Client->>Client: blocked
|
||||
Flask->>Gemini: hello
|
||||
deactivate Flask
|
||||
activate Gemini
|
||||
Flask->>Flask: blocked
|
||||
Gemini->>Flask: hi
|
||||
deactivate Gemini
|
||||
activate Flask
|
||||
Flask->>Client: hi
|
||||
deactivate Flask
|
||||
-->
|
||||
|
||||
**How Quart works:**
|
||||
|
||||
- Non-Blocking: Quart handles multiple requests concurrently. It doesn't wait for Gemini to respond before handling other messages.
|
||||
- Concurrent: The client can send messages continuously, and Quart processes them without blocking, leading to a smoother flow.
|
||||
|
||||

|
||||
|
||||
<!-- mermaid code:
|
||||
sequenceDiagram
|
||||
participant Client
|
||||
participant Quart
|
||||
participant Gemini
|
||||
|
||||
Client->>Quart: hello
|
||||
activate Client
|
||||
activate Quart
|
||||
Quart->>Gemini: hello
|
||||
activate Gemini
|
||||
|
||||
Client->>Quart: how are you?
|
||||
Gemini->>Quart: hi
|
||||
Quart->>Gemini: how are you?
|
||||
Quart->>Client: hi
|
||||
Gemini->>Quart: I'm good!
|
||||
deactivate Gemini
|
||||
Quart->>Client: I'm good!
|
||||
deactivate Quart
|
||||
deactivate Client
|
||||
-->
|
||||
|
||||
**Flask vs. Quart: Key Architectural Differences:**
|
||||
|
||||
| | | |
|
||||
| ---------------- | ----------------------------------------- | ----------------------------------------- |
|
||||
| Feature | Flask (Synchronous) | Quart (Asynchronous) |
|
||||
| Request Handling | One at a time, blocking | Concurrent, non-blocking |
|
||||
| Server Interface | WSGI | ASGI |
|
||||
| Concurrency | Through multiple processes/threads (WSGI) | Single-threaded with event loop (asyncio) |
|
||||
| View Functions | Regular def functions | async def functions |
|
||||
| I/O Operations | Blocking | Non-blocking (using await) |
|
||||
| Performance | Lower throughput for I/O-bound tasks | Higher throughput for I/O-bound tasks |
|
||||
| Complexity | Simpler to write (initially) | Steeper learning curve (async/await) |
|
||||
|
||||
### Raw WebSocket v. Quart
|
||||
|
||||
In [Gemini Multimodal Live API Demo](https://github.com/GoogleCloudPlatform/generative-ai/tree/main/gemini/multimodal-live-api/websocket-demo-app), it uses raw WebSockets API to provide a proxy function that connects the client with Gemini Live API. This is an alternative way to implement a scalable non-blocking Gen AI app with Gemini. You would typically choose this when you need maximum control, have very specific performance requirements, or are implementing a highly custom protocol.
|
||||
|
||||
Compared to it, Quart offers a higher level of abstraction, making it easier to develop, manage, and scale real-time applications built with WebSockets. It simplifies common tasks, integrates well with HTTP, and benefits from the Python ecosystem. Especially, it fit smoothly with [Google Gen AI Python SDK](https://googleapis.github.io/python-genai/index.html) and make it easier to take advantage of the high level API for handling multimodal content and function calling at the server-side.
|
||||
|
||||
## Run the demo app
|
||||
|
||||
The following sections provide instructions to run the app on Cloud Shell and deploy to Cloud Run.
|
||||
|
||||
### Download the app on Cloud Shell
|
||||
|
||||
Download the source code on [Cloud Shell](https://cloud.google.com/shell/docs/using-cloud-shell), with the following steps:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/GoogleCloudPlatform/generative-ai.git \
|
||||
gemini/sample-apps/gemini-quart-cloudrun
|
||||
cd gemini/sample-apps/gemini-quart-cloudrun
|
||||
```
|
||||
|
||||
### Run the app on Cloud Shell locally
|
||||
|
||||
To run the app on Cloud Shell locally, follow these steps:
|
||||
|
||||
1. Set project ID:
|
||||
|
||||
In Cloud Shell, execute the following commands with replacing `YOUR_PROJECT_ID`:
|
||||
|
||||
```bash
|
||||
gcloud config set project YOUR_PROJECT_ID
|
||||
```
|
||||
|
||||
1. Install the dependencies:
|
||||
|
||||
```bash
|
||||
pip install -r app/requirements.txt
|
||||
```
|
||||
|
||||
1. To run the app locally, execute the following command:
|
||||
|
||||
```bash
|
||||
cd app
|
||||
chmod +x run.sh
|
||||
./run.sh
|
||||
```
|
||||
|
||||
1. (Optional) To run the app with Gemini API key:
|
||||
|
||||
If you like to run the app with Gemini API key instead of Vertex AI, edit `run.sh` to specify your [Gemini API Key](https://aistudio.google.com/apikey).
|
||||
|
||||
The application will start up. Use Cloud Shell's [web preview](https://cloud.google.com/shell/docs/using-web-preview) button at top right to launch the preview page. You may also visit that in the browser to view the application.
|
||||
|
||||
## Build and Deploy the Application to Cloud Run
|
||||
|
||||
To deploy the Quart Application in [Cloud Run](https://cloud.google.com/run/docs/quickstarts/deploy-container), we need to perform the following steps:
|
||||
|
||||
1. Set project ID:
|
||||
|
||||
In Cloud Shell, execute the following commands with replacing `YOUR_PROJECT_ID`:
|
||||
|
||||
```bash
|
||||
gcloud config set project YOUR_PROJECT_ID
|
||||
```
|
||||
|
||||
1. To deploy the app to Cloud Run, execute the following command:
|
||||
|
||||
```bash
|
||||
cd app
|
||||
chmod +x deploy.sh
|
||||
./deploy.sh
|
||||
```
|
||||
|
||||
On successful deployment, you will be provided a URL to the Cloud Run service. You can visit that in the browser to view the Cloud Run application that you just deployed.
|
||||
|
||||
### If you see `RESOURCE_EXHAUSTED` errors
|
||||
|
||||
While running the app using Vertex AI, you might occasionally encounter `RESOURCE_EXHAUSTED` errors on the Cloud Run logs tab. This typically means you've hit the quota limit on the number of concurrent sessions you can open with the Gemini API. If this happens, you have a couple of options: you can either wait a few minutes and try running the app again, or switch to using the Gemini Developer API by specifying your [Gemini API Key](https://aistudio.google.com/apikey) in the `run.sh` or `deploy.sh` script accordintly. This can provide a workaround.
|
||||
|
||||
Congratulations!
|
||||
|
||||
## How the demo app works
|
||||
|
||||
### How `app.py` works
|
||||
|
||||
The `app.py` file defines a Quart web application that facilitates real-time interaction with the Google Gemini API for large language model processing. Here's a breakdown of the flow:
|
||||
|
||||
- **WebSocket Endpoint (`/live`):** The /live route establishes a WebSocket connection for real-time communication with Gemini. This is the core of the application's interactive functionality.
|
||||
|
||||
- **WebSocket Handlers (`upstream_worker` and `downstream_worker`):** Within the `/live` WebSocket handler, two asynchronous tasks are created:
|
||||
|
||||
- **upstream_worker:** This task continuously reads messages from the client's WebSocket connection and sends them to the Gemini API using `gemini_session.send()`. Each message from the client is treated as a turn in the conversation.
|
||||
|
||||
- **downstream_worker:** This task continuously receives streaming responses from Gemini using `gemini_session.receive()`. It then formats these responses into JSON packets containing the text and turn completion status, and sends them back to the client via the WebSocket.
|
||||
|
||||
- **Concurrency Management:** The `upstream_worker` and `downstream_worker` operate concurrently using `asyncio`. This enables bidirectional, real-time communication between the client and Gemini. The `asyncio.wait()` function is used to monitor both tasks for exceptions, allowing the application to handle errors gracefully.
|
||||
|
||||
- **Session Management:** The `gemini_session` is established within an async with block, ensuring that the session is properly closed when the WebSocket connection is terminated. This prevents resource leaks and maintains a clean state.
|
||||
|
||||
### How `index.html` works
|
||||
|
||||
- **Structure of `index.html`**: The HTML sets up a basic page with a title, a heading ("Gemini Live API Test"), a message display area (messages div), and a form for sending messages.
|
||||
|
||||
- **WebSocket Connection:** The core functionality lies in the JavaScript section. It establishes a WebSocket connection to the `/live` endpoint on the same host as the page.
|
||||
|
||||
- **WebSocket Event Handlers:** Several event handlers manage the WebSocket interaction:
|
||||
|
||||
- **`onopen`:** When the WebSocket connection is successfully established, this handler enables the `Send` button, displays a `Connection opened` message, and adds a submit handler to the message form.
|
||||
|
||||
- **`onmessage`:** This handler processes incoming messages from the server (Gemini responses). It parses the JSON data, checks for turn completion, updates message display with response, scrolls messages into view, creates new message entry for new turns, and displays ongoing responses piece by piece for incomplete turns.
|
||||
|
||||
- **`onclose`:** This handler is called when the WebSocket connection is closed. It disables the `Send` button, displays a `Connection closed` message, and initiates a timer to retry connecting to the server in 5 seconds.
|
||||
|
||||
### Improvement for production deployment
|
||||
|
||||
While this is a minimal demo app, you could extend it to a production app by improving the following areas:
|
||||
|
||||
- **Handling audio and images:** The application can be extended to support audio and images. See [Getting Started with the Multimodal Live API using Gen AI SDK](https://github.com/GoogleCloudPlatform/generative-ai/blob/main/gemini/multimodal-live-api/intro_multimodal_live_api_genai_sdk.ipynb) on how to process the multimodal content.
|
||||
|
||||
- **Gemini Live API Rate Limits:** The application doesn't handle [the Gemini Live API rate limits](https://ai.google.dev/api/multimodal-live#rate-limits). In production you need a rate throttling mechanism for the `concurrent sessions per key` and `tokens per minute` to handle traffic from multiple clients.
|
||||
|
||||
- **Security:** The `allow-unauthenticated` flag in `deploy.sh` makes the application publicly accessible. For production use, authentication and authorization should be implemented to control access.
|
||||
|
||||
- **Session Management:** While the current session management within the WebSocket handler is functional, more robust session handling could be explored for scenarios involving multiple users or persistent sessions.
|
||||
|
||||
## References
|
||||
|
||||
- [Gemini Multimodal Live API](https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/multimodal-live)
|
||||
- [Gemini Multimodal Live API Demo](https://github.com/GoogleCloudPlatform/generative-ai/tree/main/gemini/multimodal-live-api/websocket-demo-app)
|
||||
- [Google Gen AI Python SDK](https://googleapis.github.io/python-genai/index.html)
|
||||
- [Getting Started with the Multimodal Live API using Gen AI SDK](https://github.com/GoogleCloudPlatform/generative-ai/blob/main/gemini/multimodal-live-api/intro_multimodal_live_api_genai_sdk.ipynb)
|
||||
- [Quart documents](https://quart.palletsprojects.com/en/latest/)
|
||||
@@ -0,0 +1,18 @@
|
||||
# Use a slim Python base image
|
||||
FROM python:3.13-slim
|
||||
|
||||
# Set working directory
|
||||
WORKDIR /app
|
||||
|
||||
# Copy dependencies
|
||||
COPY requirements.txt requirements.txt
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
# Copy your app code
|
||||
COPY . .
|
||||
|
||||
# Expose port 8080
|
||||
EXPOSE 8080
|
||||
|
||||
# Run hypercorn
|
||||
CMD ["hypercorn", "app:app", "--bind", "0.0.0.0:8080"]
|
||||
@@ -0,0 +1,151 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# pylint: disable=too-many-lines
|
||||
# pylint: disable=import-error
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from typing import Any, Dict
|
||||
|
||||
from google.genai import Client
|
||||
from google.genai.live import AsyncSession
|
||||
from google.genai.types import LiveConnectConfig
|
||||
from quart import Quart, Response, Websocket, send_from_directory, websocket
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
|
||||
#
|
||||
# Gemini API
|
||||
#
|
||||
|
||||
PROJECT_ID: str = os.environ.get("PROJECT_ID", "")
|
||||
LOCATION: str = os.environ.get("LOCATION", "us-central1")
|
||||
GEMINI_API_KEY: str = os.environ.get("GEMINI_API_KEY", "")
|
||||
QUART_DEBUG_MODE: bool = os.environ.get("QUART_DEBUG_MODE") == "True"
|
||||
|
||||
GEMINI_MODEL: str = "gemini-2.0-flash-live-preview-04-09"
|
||||
|
||||
# Gemini API Client: Use either one of the following APIs
|
||||
gemini_client: Client = (
|
||||
Client(vertexai=True, project=PROJECT_ID, location=LOCATION)
|
||||
if not GEMINI_API_KEY
|
||||
else Client(api_key=GEMINI_API_KEY, http_options={"api_version": "v1alpha"})
|
||||
)
|
||||
|
||||
# Gemini API config
|
||||
gemini_config = LiveConnectConfig(
|
||||
response_modalities=["TEXT"],
|
||||
)
|
||||
|
||||
#
|
||||
# Quart
|
||||
#
|
||||
|
||||
app: Quart = Quart(__name__)
|
||||
|
||||
|
||||
@app.route("/")
|
||||
async def index() -> Response:
|
||||
"""
|
||||
Serve index.html for the index access.
|
||||
"""
|
||||
return await send_from_directory("public", "index.html")
|
||||
|
||||
|
||||
async def upstream_worker(
|
||||
gemini_session: AsyncSession, client_websocket: Websocket
|
||||
) -> None:
|
||||
"""
|
||||
Continuously read messages from the client WebSocket
|
||||
and forward them to Gemini.
|
||||
"""
|
||||
while True:
|
||||
message: str = await client_websocket.receive()
|
||||
await gemini_session.send(input=message, end_of_turn=True)
|
||||
logging.info(
|
||||
"upstream_worker(): sent a message from client to Gemini: %s", message
|
||||
)
|
||||
|
||||
|
||||
async def downstream_worker(
|
||||
gemini_session: AsyncSession, client_websocket: Websocket
|
||||
) -> None:
|
||||
"""
|
||||
Continuously read streaming responses from Gemini
|
||||
and send them directly to the client WebSocket.
|
||||
"""
|
||||
while True:
|
||||
async for response in gemini_session.receive():
|
||||
if not response:
|
||||
continue
|
||||
|
||||
packet: Dict[str, Any] = {
|
||||
"text": response.text if response.text else "",
|
||||
"turn_complete": response.server_content.turn_complete,
|
||||
}
|
||||
await client_websocket.send(json.dumps(packet))
|
||||
logging.info("downstream_worker(): sent response to client: %s", packet)
|
||||
|
||||
|
||||
@app.websocket("/live")
|
||||
async def live() -> None:
|
||||
"""
|
||||
WebSocket endpoint for live (streaming) connections to Gemini.
|
||||
"""
|
||||
|
||||
# Connect to Gemini in "live" (streaming) mode
|
||||
async with gemini_client.aio.live.connect(
|
||||
model=GEMINI_MODEL, config=gemini_config
|
||||
) as gemini_session:
|
||||
upstream_task: asyncio.Task = asyncio.create_task(
|
||||
upstream_worker(gemini_session, websocket)
|
||||
)
|
||||
downstream_task: asyncio.Task = asyncio.create_task(
|
||||
downstream_worker(gemini_session, websocket)
|
||||
)
|
||||
logging.info("live(): connected to Gemini, started workers.")
|
||||
|
||||
try:
|
||||
# Wait until either task finishes or raises an exception
|
||||
done, pending = await asyncio.wait(
|
||||
[downstream_task, upstream_task], return_when=asyncio.FIRST_EXCEPTION
|
||||
)
|
||||
|
||||
# If one of them raised, re-raise that exception here
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
for task in done:
|
||||
exc = task.exception()
|
||||
if exc:
|
||||
raise exc
|
||||
|
||||
# Handle cancelled errors
|
||||
except asyncio.CancelledError:
|
||||
logging.info("live(): client connection closed.")
|
||||
|
||||
finally:
|
||||
# Cancel any leftover tasks
|
||||
upstream_task.cancel()
|
||||
downstream_task.cancel()
|
||||
await asyncio.gather(downstream_task, upstream_task, return_exceptions=True)
|
||||
|
||||
# Close Gemini session
|
||||
await gemini_session.close()
|
||||
logging.info("live(): Gemini session closed.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
app.run(host="0.0.0.0", port=8080, debug=QUART_DEBUG_MODE)
|
||||
+32
@@ -0,0 +1,32 @@
|
||||
#!/bin/bash
|
||||
#set -e
|
||||
|
||||
# To run the app with Vertex AI, use this script as is.
|
||||
# To run the app with Gemini API key, comment out the following lines.
|
||||
PROJECT_ID=$(gcloud config get-value project)
|
||||
export PROJECT_ID
|
||||
LOCATION=us-central1
|
||||
export LOCATION
|
||||
|
||||
# To run the app with Gemini API key, uncomment this and specify your key.
|
||||
# (See: https://aistudio.google.com/apikey)
|
||||
#export GEMINI_API_KEY=<YOUR GEMINI API KEY>
|
||||
|
||||
# Quart debug mode (True or False)
|
||||
QUART_DEBUG_MODE=False
|
||||
export QUART_DEBUG_MODE
|
||||
|
||||
# build an image
|
||||
gcr_image_path=gcr.io/$PROJECT_ID/gemini-quart-cloudrun
|
||||
gcloud builds submit --tag $gcr_image_path
|
||||
|
||||
# deploy
|
||||
gcloud run deploy gemini-quart-cloudrun \
|
||||
--image $gcr_image_path \
|
||||
--platform managed \
|
||||
--allow-unauthenticated \
|
||||
--project=$PROJECT_ID --region=$LOCATION \
|
||||
--set-env-vars=PROJECT_ID=$PROJECT_ID \
|
||||
--set-env-vars=LOCATION=$LOCATION \
|
||||
--set-env-vars=GEMINI_API_KEY=$GEMINI_API_KEY \
|
||||
--set-env-vars=QUART_DEBUG_MODE=$QUART_DEBUG_MODE
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user