In [2]:
from langchain.agents.mrkl.base import ZeroShotAgent
from langchain.llms.openai import OpenAI
import os
from langchain.sql_database import SQLDatabase
from langchain.agents.agent_toolkits import SQLDatabaseToolkit

from langchain.agents.agent_toolkits.sql.prompt import (
    SQL_FUNCTIONS_SUFFIX,
    SQL_PREFIX,
    SQL_SUFFIX,
)

from langchain.agents.mrkl.prompt import FORMAT_INSTRUCTIONS
from langchain.agents.agent import AgentExecutor
from langchain.chains.llm import LLMChain

# I. Step-by-Step Guide to Creating an SQL Agent

![image.png](attachment:4433946e-916f-4d11-9abd-de3d8500a969.png)

**Step 1:** Create the Language Model

In [9]:
model = OpenAI(temperature=0, openai_api_key=os.getenv("OPENAI_API"))

**Step 2:** Define Connection and SQL Database Tools


In [10]:
db = SQLDatabase.from_uri("teradatasql://dbc:dbc@192.168.11.7:1025/Sales")
toolkit = SQLDatabaseToolkit(db=db, llm=model)
tools = toolkit.get_tools()

**Step 3:** Define the Prompt

In [11]:
# Define the prompt with relevant settings
prefix = SQL_PREFIX.format(dialect=toolkit.dialect, top_k=10)

prompt = ZeroShotAgent.create_prompt(
            tools,
            prefix=prefix,
            suffix=SQL_SUFFIX,
            format_instructions=FORMAT_INSTRUCTIONS,
            input_variables=None,
        )

**Step 4:** Define the Language Model Chain

In [12]:
# Step 4. Define Chain
llm_chain = LLMChain(
    llm=model,
    prompt=prompt,
    callback_manager=None,
)

In [13]:
# Step 5. Define Agent

agent = ZeroShotAgent(
    llm_chain=llm_chain,
    prompt=prompt,
    tools=tools,
    callback_manager=None,
)

In [14]:
# Step 6. Define Agent Executor
agent_executor = AgentExecutor.from_agent_and_tools(
    agent=agent,
    tools=tools,
    callback_manager=None,
    verbose=True,
)

In [27]:
agent_executor.run("Count the customer")



[1m> Entering new AgentExecutor chain...[0m
[32;1m[1;3mAction: sql_db_list_tables
Action Input: [0m
Observation: [38;5;200m[1;3mcustomer, orderdetail, product[0m
Thought:[32;1m[1;3m I should query the schema of the customer table
Action: sql_db_schema
Action Input: customer[0m
Observation: [33;1m[1;3m
CREATE TABLE customer (
	customer_id VARCHAR(255) CHAR SET LATIN NOT NULL, 
	customer_name VARCHAR(255) CHAR SET LATIN NOT NULL, 
	customer_email VARCHAR(255) CHAR SET LATIN NOT NULL, 
	customer_address VARCHAR(255) CHAR SET LATIN NOT NULL, 
	PRIMARY KEY (customer_id)
)

/*
3 rows from customer table:
customer_id	customer_name	customer_email	customer_address
   8	Grace Davis	grace@example.com	555 Maple St
   4	Alice Johnson	alice@example.com	111 Pine St
   3	Bob Johnson	bob@example.com	789 Pine St
*/[0m
Thought:[32;1m[1;3m I should query the customer table to count the customers
Action: sql_db_query
Action Input: SELECT COUNT(*) FROM customer[0m
Observation: [36;1m[1;

'There are 10 customers.'

# II. The Inner Workings of the SQL Agent

![image.png](attachment:d1f68935-02f3-4371-b6ab-331d41c54cbb.png)

**Common function, init**

In [15]:
def construct_scratchpad(intermediate_steps):
    """Construct the scratchpad that lets the agent continue its thought process."""
    observation_prefix = 'Observation: '
    llm_prefix = 'Thought:'
    
    thoughts = ""
    for action, observation in intermediate_steps:
        try:
            thoughts += action.log
        except Exception as e:
            thoughts = ""
        thoughts += f"\n{observation_prefix}{observation}\n{llm_prefix}"
    return thoughts

def prep_prompts(input_prompt, input_list):
    """Prepare prompts from inputs."""
    
    stop = None
    if len(input_list) == 0:
        return [], stop
    if "stop" in input_list[0]:
        stop = input_list[0]["stop"]
    prompts = []
    for inputs in input_list:
        selected_inputs = {k: inputs[k] for k in input_prompt.input_variables}
        prompt = input_prompt.format_prompt(**selected_inputs)
        _colored_text = "green"
        _text = "Prompt after formatting:\n" + _colored_text
        prompts.append(prompt)
    return prompts, stop

name_to_tool_map = {tool.name: tool for tool in agent_executor.tools}

## Input Question

In [16]:
question = "Count the customer"

In [17]:
intermediate_steps = []
input = {'input': question}

## LOOP

### Round 1

#### 1. Prompt

In [18]:
# get_full_inputs function 
# Define though input. First loop, thoughts = None. 
thoughts = construct_scratchpad(intermediate_steps)
new_inputs = {"agent_scratchpad": thoughts, "stop": ['\nObservation:', '\n\tObservation:']}

# Full input 
full_inputs = {**input, **new_inputs}

# LLM chain, define promot
prompts, stop = prep_prompts(prompt, [full_inputs])

In [None]:
print(prompts[0].text)

#### 2. Model

In [20]:
# Chain will call following function agent.plan() ->  full_inputs = self.get_full_inputs(intermediate_steps, **kwargs) full_output = llm_chain.predict(**full_inputs)
model_output = model.generate_prompt(prompts, stop)
full_output = model_output.generations[0][0].text

In [21]:
full_output

'Action: sql_db_list_tables\nAction Input: '

#### 3. Output Parser

Agent will parse output from model to AgentAction

In [22]:
# Parse full output to Action or Finish Agent
agent_output = agent.output_parser.parse(full_output)

In [23]:
agent_output

AgentAction(tool='sql_db_list_tables', tool_input='', log='Action: sql_db_list_tables\nAction Input: ')

#### 4. Tools 
**Run an action**

In [24]:
# Get tool 
exe_tool = name_to_tool_map.get(agent_output.tool)

In [25]:
kwargs = agent.tool_run_logging_kwargs()

observation = exe_tool.run(agent_output.tool_input,
                      **kwargs)

result = [(exe_tool, observation)]

In [26]:
intermediate_steps.extend(result)

In [27]:
observation

'customer, orderdetail, product, table1, table2'

### Round 2

#### 1. Prompt

In [28]:
# Define though input. First loop, thoughts = None. 
thoughts = construct_scratchpad(intermediate_steps)
new_inputs = {"agent_scratchpad": thoughts, "stop": ['\nObservation:', '\n\tObservation:']}

# Full input 
full_inputs = {**input, **new_inputs}

# LLM chain, define promot
prompts, stop = prep_prompts(prompt, [full_inputs])

In [None]:
print(prompts[0].text)

#### 2. Model

In [30]:
# Chain will call following function agent.plan() ->  full_inputs = self.get_full_inputs(intermediate_steps, **kwargs) full_output = llm_chain.predict(**full_inputs)
model_output = model.generate_prompt(prompts, stop)
full_output = model_output.generations[0][0].text

In [31]:
full_output

' The customer table is the most relevant, so I should query the schema of that table.\n\nAction: sql_db_schema\nAction Input: customer'

#### 3. Output Parser

In [32]:
# Parse full output to Action or Finish Agent
agent_output = agent.output_parser.parse(full_output)

In [33]:
agent_output

AgentAction(tool='sql_db_schema', tool_input='customer', log=' The customer table is the most relevant, so I should query the schema of that table.\n\nAction: sql_db_schema\nAction Input: customer')

#### 4. Tools

In [34]:
exe_tool = name_to_tool_map.get(agent_output.tool)

In [35]:
exe_tool

InfoSQLDatabaseTool(description="Input to this tool is a comma-separated list of tables, output is the schema and sample rows for those tables. Be sure that the tables actually exist by calling sql_db_list_tables first! Example Input: 'table1, table2, table3'", db=<langchain.utilities.sql_database.SQLDatabase object at 0x000002526A7374C0>)

In [36]:
kwargs = agent.tool_run_logging_kwargs()

observation = exe_tool.run(agent_output.tool_input,
                      **kwargs)

result = [(exe_tool, observation)]

In [None]:
result

In [38]:
intermediate_steps.extend(result)

In [39]:
print(observation)


CREATE TABLE customer (
	customer_id VARCHAR(255) CHAR SET LATIN NOT NULL, 
	customer_name VARCHAR(255) CHAR SET LATIN NOT NULL, 
	customer_email VARCHAR(255) CHAR SET LATIN NOT NULL, 
	customer_address VARCHAR(255) CHAR SET LATIN NOT NULL, 
	PRIMARY KEY (customer_id)
)

/*
3 rows from customer table:
customer_id	customer_name	customer_email	customer_address
   8	Grace Davis	grace@example.com	555 Maple St
   4	Alice Johnson	alice@example.com	111 Pine St
   3	Bob Johnson	bob@example.com	789 Pine St
*/


### Round 3

#### 1. Prompt

In [40]:
# Define though input. First loop, thoughts = None. 
thoughts = construct_scratchpad(intermediate_steps)
new_inputs = {"agent_scratchpad": thoughts, "stop": ['\nObservation:', '\n\tObservation:']}

# Full input 
full_inputs = {**input, **new_inputs}

# LLM chain, define promot
prompts, stop = prep_prompts(prompt, [full_inputs])

In [None]:
print(prompts[0].text)

#### 2. Model

In [45]:
# Chain will call following function agent.plan() ->  full_inputs = self.get_full_inputs(intermediate_steps, **kwargs) full_output = llm_chain.predict(**full_inputs)
model_output = model.generate_prompt(prompts, stop)
full_output = model_output.generations[0][0].text

In [46]:
full_output

' I should query the customer table to count the customers\nAction: sql_db_query\nAction Input: SELECT COUNT(*) FROM customer'

#### 3. Output Parser

In [47]:
# Parse full output to Action or Finish Agent
agent_output = agent.output_parser.parse(full_output)

In [48]:
agent_output

AgentAction(tool='sql_db_query', tool_input='SELECT COUNT(*) FROM customer', log=' I should query the customer table to count the customers\nAction: sql_db_query\nAction Input: SELECT COUNT(*) FROM customer')

#### 4. Tools

In [49]:
exe_tool = name_to_tool_map.get(agent_output.tool)

In [50]:
kwargs = agent.tool_run_logging_kwargs()

observation = exe_tool.run(agent_output.tool_input,
                      **kwargs)

result = [(exe_tool, observation)]

In [51]:
result

[(QuerySQLDataBaseTool(description="Input to this tool is a detailed and correct SQL query, output is a result from the database. If the query is not correct, an error message will be returned. If an error is returned, rewrite the query, check the query, and try again. If you encounter an issue with Unknown column 'xxxx' in 'field list', use sql_db_schema to query the correct table fields.", db=<langchain.utilities.sql_database.SQLDatabase object at 0x00000126FD1055E0>),
  '[(10,)]')]

In [52]:
intermediate_steps.extend(result)

In [53]:
observation

'[(10,)]'

### Round 4

#### 1. Prompt

In [54]:
# Define though input. First loop, thoughts = None. 
thoughts = construct_scratchpad(intermediate_steps)
new_inputs = {"agent_scratchpad": thoughts, "stop": ['\nObservation:', '\n\tObservation:']}

# Full input 
full_inputs = {**input, **new_inputs}

# LLM chain, define promot
prompts, stop = prep_prompts(prompt, [full_inputs])

In [None]:
print(prompts[0].text)

#### 2. Model

In [56]:
# Chain will call following function agent.plan() ->  full_inputs = self.get_full_inputs(intermediate_steps, **kwargs) full_output = llm_chain.predict(**full_inputs)
model_output = model.generate_prompt(prompts, stop)
full_output = model_output.generations[0][0].text

In [57]:
full_output

' I now know the final answer\nFinal Answer: There are 10 customers.'

#### 3. Output Parser

In [58]:
# Parse full output to Action or Finish Agent
agent_output = agent.output_parser.parse(full_output)

In [59]:
agent_output

AgentFinish(return_values={'output': 'There are 10 customers.'}, log=' I now know the final answer\nFinal Answer: There are 10 customers.')

### END LOOP

At round 4, after parse output from 