Team Ai
Apppublic

vedsadani/Text2SQL

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
app.py136 linesDownload Raw Back to root
1import pandas as pd2from openai import OpenAI3import os4from google.cloud import bigquery5import numpy as np6import gradio as gr7 8project_id = os.getenv('project_id')9dataset_id = os.getenv('dataset_id')10table_id = os.getenv('table_id')11 12openai_client = OpenAI()13 14def fetch_table_schema(project_id, dataset_id, table_id):15    bqclient = bigquery.Client(project=project_id)16 17    table_ref = f"{project_id}.{dataset_id}.{table_id}"18 19    table = bqclient.get_table(table_ref)20 21    schema_dict = {}22    for schema_field in table.schema:23        schema_dict[schema_field.name] = schema_field.field_type24 25    return schema_dict26 27def get_sql_query(description):28    prompt = f'''29    Generate the SQL query for the following task:\n{description}.\n30    The database you need is called {dataset_id} and the table is called {table_id}.31    Use the format {dataset_id}.{table_id} as the table name in the queries.32    Enclose column names in backticks(`) not quotation marks.33    Do not assign aliases to the columns.34    Do not calculate new columns, unless specifically called to.35    Return only the SQL query, nothing else.36    Do not use WITHIN GROUP clause.37    \nThe list of all the columns is as follows: {schema} /n38    '''39    try:40      completion = openai_client.chat.completions.create(41          model='gpt-4o',42          messages = [43                  {"role": "system", "content": "You are an expert Data Scientist with in-depth knowledge of SQL, working on Network Telemetry Data."},44                  {"role": "user", "content": f'{prompt}'},45                ]46      )47      sql_query = completion.choices[0].message.content.strip().split('```sql')[1].split('```')[0]48 49    except Exception as e:50     print(f'The following error ocurred: {e}\n')51     sql_query = None52 53    return sql_query54 55schema = fetch_table_schema(project_id, dataset_id, table_id)56 57def execute_sql_query(query):58    client = bigquery.Client()59 60    try:61     result = client.query(query).to_dataframe()62     message = f'The query:{query} was successfully executed.'63 64    except Exception as e:65     result = None66     message = f'The query:{query} could not be executed due to the following exception:\n{e}'67 68    return result, message69 70def echo(text):71  query = get_sql_query(text)72  if query is None:73    return 'No query generated', 'No query generated'74  result, message = execute_sql_query(query)75  return result, message76 77def gradio_interface(text):78    result, message = echo(text)79    if isinstance(result, pd.DataFrame):80        return gr.Dataframe(value=result), message81    else:82        return result, message83 84def gradio_interface(text):85    result, message = echo(text)86    if isinstance(result, pd.DataFrame):87        return gr.Dataframe(value=result), message88    else:89        return result, message90 91demo = gr.Blocks(92        title="Text-to-SQL",93        theme='remilia/ghostly',94)95 96with demo:97 98  gr.Markdown(99    '''100    # <p style="text-align: center;">Text to SQL Query Engine</p>101 102    <p style="text-align: center;">103    Welcome to our Text2SQL Engine.104    <br>105    Enter your query in natural language and we'll convert it to SQL and return the result to you.106    </p>107    '''108    )109 110  with gr.Row():111    with gr.Column(scale=1):112      text_input = gr.Textbox(label="Enter your query")113      button = gr.Button("Submit")114      gr.Examples([115        'Find the correlation between RTT and Jitter for each Market',116        'Find the variance in Jitter for each 5G_Reliability_Category',117        'Find the count of records per 5G_Reliability_Category where 5G_Reliability_Value is below the average for the category',118        'Calculate the standard deviation of 5G_Reliability_Score for each Network_Engineer',119        'Determine the Sector with the highest variance in 5G Reliability Value and its corresponding average Context Drop Percent'120        ],121                inputs=[text_input]122                )123    with gr.Column(scale=3):124      output_text = gr.Textbox(label="Output", interactive=False)125      output_df = gr.Dataframe(interactive=False)126 127    def update_output(text):128        result, message = gradio_interface(text)129        if result and isinstance(result, pd.DataFrame):130            return result, message, gr.update(visible=True)131        else:132            return result, message, gr.update(visible=False)133 134    button.click(update_output, inputs=text_input, outputs=[output_df, output_text])135 136demo.launch(debug=True, auth=("admin", "Text2SQL"))