# SQL Database Agent

This notebook showcases an agent designed to interact with a sql databases. The agent builds off of [SQLDatabaseChain](https://langchain.readthedocs.io/en/latest/modules/chains/examples/sqlite.html) and is designed to answer more general questions about a database, as well as recover from errors.

Note that, as this agent is in active development, all answers might not be correct. Additionally, it is not guaranteed that the agent won't perform DML statements on your database given certain questions. Be careful running it on sensitive data!

This uses the example Chinook database. To set it up follow the instructions on https://database.guide/2-sample-databases-sqlite/, placing the .db file in a notebooks folder at the root of this repository.

## Initialization

In [21]:
from langchain.agents import create_sql_agent
from langchain.agents.agent_toolkits import SQLDatabaseToolkit
from langchain.sql_database import SQLDatabase
from langchain.llms.openai import OpenAI
from langchain.agents import AgentExecutor

In [22]:
db = SQLDatabase.from_uri("sqlite:///../../../../../notebooks/Chinook.db")
llm = OpenAI(temperature=0)
toolkit = SQLDatabaseToolkit(db=db, llm=llm)
agent_executor = create_sql_agent(
    llm=llm,
    toolkit=toolkit,
    verbose=True
)

## Example: describing a table

In [23]:
agent_executor.run("Describe the playlisttrack table")

Retrying langchain.llms.openai.completion_with_retry.<locals>._completion_with_retry in 4.0 seconds as it raised APIConnectionError: Error communicating with OpenAI: ('Connection aborted.', ConnectionResetError(54, 'Connection reset by peer')).




[1m> Entering new AgentExecutor chain...[0m
[32;1m[1;3mAction: list_tables_sql_db
Action Input: ""[0m
Observation: [38;5;200m[1;3mInvoice, MediaType, Genre, Album, Customer, InvoiceLine, PlaylistTrack, Playlist, Employee, Artist, Track[0m
Thought:[32;1m[1;3m I should look at the schema of the playlisttrack table
Action: schema_sql_db
Action Input: "PlaylistTrack"[0m
Observation: [33;1m[1;3m
CREATE TABLE "PlaylistTrack" (
	"PlaylistId" INTEGER NOT NULL, 
	"TrackId" INTEGER NOT NULL, 
	PRIMARY KEY ("PlaylistId", "TrackId"), 
	FOREIGN KEY("TrackId") REFERENCES "Track" ("TrackId"), 
	FOREIGN KEY("PlaylistId") REFERENCES "Playlist" ("PlaylistId")
)

/*
3 rows from PlaylistTrack table:
PlaylistId	TrackId
1	3402
1	3389
1	3390
*/[0m
Thought:[32;1m[1;3m I now know the final answer
Final Answer: The PlaylistTrack table has two columns, PlaylistId and TrackId, and is related to the Playlist and Track tables.[0m

[1m> Finished chain.[0m


'The PlaylistTrack table has two columns, PlaylistId and TrackId, and is related to the Playlist and Track tables.'

## Example: describing a table, recovering from an error

In this example, the agent tries to search for a table that doesn't exist, but finds the next best result

In [24]:
agent_executor.run("Describe the playlistsong table")



[1m> Entering new AgentExecutor chain...[0m
[32;1m[1;3mAction: list_tables_sql_db
Action Input: ""[0m
Observation: [38;5;200m[1;3mInvoice, MediaType, Genre, Album, Customer, InvoiceLine, PlaylistTrack, Playlist, Employee, Artist, Track[0m
Thought:[32;1m[1;3m I should look at the schema of the playlistsong table
Action: schema_sql_db
Action Input: "PlaylistTrack"[0m
Observation: [33;1m[1;3m
CREATE TABLE "PlaylistTrack" (
	"PlaylistId" INTEGER NOT NULL, 
	"TrackId" INTEGER NOT NULL, 
	PRIMARY KEY ("PlaylistId", "TrackId"), 
	FOREIGN KEY("TrackId") REFERENCES "Track" ("TrackId"), 
	FOREIGN KEY("PlaylistId") REFERENCES "Playlist" ("PlaylistId")
)

/*
3 rows from PlaylistTrack table:
PlaylistId	TrackId
1	3402
1	3389
1	3390
*/[0m
Thought:[32;1m[1;3m I now know the final answer
Final Answer: The PlaylistTrack table has two columns, PlaylistId and TrackId, and is used to store the relationship between playlists and tracks.[0m

[1m> Finished chain.[0m


'The PlaylistTrack table has two columns, PlaylistId and TrackId, and is used to store the relationship between playlists and tracks.'

## Example: running queries

In [25]:
agent_executor.run("List the total sales per country. Which country's customers spent the most?")



[1m> Entering new AgentExecutor chain...[0m
[32;1m[1;3mAction: list_tables_sql_db
Action Input: ""[0m
Observation: [38;5;200m[1;3mInvoice, MediaType, Genre, Album, Customer, InvoiceLine, PlaylistTrack, Playlist, Employee, Artist, Track[0m
Thought:[32;1m[1;3m I should look at the schema of the relevant tables to see what columns I can use.
Action: schema_sql_db
Action Input: "Invoice, Customer"[0m
Observation: [33;1m[1;3m
CREATE TABLE "Customer" (
	"CustomerId" INTEGER NOT NULL, 
	"FirstName" NVARCHAR(40) NOT NULL, 
	"LastName" NVARCHAR(20) NOT NULL, 
	"Company" NVARCHAR(80), 
	"Address" NVARCHAR(70), 
	"City" NVARCHAR(40), 
	"State" NVARCHAR(40), 
	"Country" NVARCHAR(40), 
	"PostalCode" NVARCHAR(10), 
	"Phone" NVARCHAR(24), 
	"Fax" NVARCHAR(24), 
	"Email" NVARCHAR(60) NOT NULL, 
	"SupportRepId" INTEGER, 
	PRIMARY KEY ("CustomerId"), 
	FOREIGN KEY("SupportRepId") REFERENCES "Employee" ("EmployeeId")
)

/*
3 rows from Customer table:
CustomerId	FirstName	LastName	Company	A

'The customers from the USA spent the most, with a total of $523.06.'

In [7]:
agent_executor.run("Show the total number of tracks in each playlist. The Playlist name should be included in the result.")



[1m> Entering new AgentExecutor chain...[0m
[32;1m[1;3mAction: list_tables_sql_db
Action Input: ""[0m
Observation: [38;5;200m[1;3mInvoice, MediaType, Artist, InvoiceLine, Genre, Playlist, Employee, Album, PlaylistTrack, Track, Customer[0m
Thought:[32;1m[1;3m I should look at the schema of the Playlist and PlaylistTrack tables to see what columns I can use.
Action: schema_sql_db
Action Input: "Playlist, PlaylistTrack"[0m
Observation: [33;1m[1;3m
CREATE TABLE "Playlist" (
	"PlaylistId" INTEGER NOT NULL, 
	"Name" NVARCHAR(120), 
	PRIMARY KEY ("PlaylistId")
)

SELECT * FROM 'Playlist' LIMIT 3;
PlaylistId Name
1 Music
2 Movies
3 TV Shows


CREATE TABLE "PlaylistTrack" (
	"PlaylistId" INTEGER NOT NULL, 
	"TrackId" INTEGER NOT NULL, 
	PRIMARY KEY ("PlaylistId", "TrackId"), 
	FOREIGN KEY("TrackId") REFERENCES "Track" ("TrackId"), 
	FOREIGN KEY("PlaylistId") REFERENCES "Playlist" ("PlaylistId")
)

SELECT * FROM 'PlaylistTrack' LIMIT 3;
PlaylistId TrackId
1 3402
1 3389
1 3390[0m


"The total number of tracks in each playlist are: '90’s Music' (1477), 'Brazilian Music' (39), 'Classical' (75), 'Classical 101 - Deep Cuts' (25), 'Classical 101 - Next Steps' (25), 'Classical 101 - The Basics' (25), 'Grunge' (15), 'Heavy Metal Classic' (26), 'Music' (6580), 'Music Videos' (1)."

## Recovering from an error

In this example, the agent is able to recover from an error after initially trying to access an attribute (`Track.ArtistId`) which doesn't exist.

In [16]:
agent_executor.run("Who are the top 3 best selling artists?")



[1m> Entering new AgentExecutor chain...[0m
[32;1m[1;3mAction: list_tables_sql_db
Action Input: ""[0m
Observation: [38;5;200m[1;3mMediaType, Track, Invoice, Album, Playlist, Customer, Employee, InvoiceLine, PlaylistTrack, Genre, Artist[0m
Thought:[32;1m[1;3m I should look at the schema of the Artist, InvoiceLine, and Track tables to see what columns I can use.
Action: schema_sql_db
Action Input: "Artist, InvoiceLine, Track"[0m
Observation: [33;1m[1;3m
CREATE TABLE "Artist" (
	"ArtistId" INTEGER NOT NULL, 
	"Name" NVARCHAR(120), 
	PRIMARY KEY ("ArtistId")
)

SELECT * FROM 'Artist' LIMIT 3;
ArtistId Name
1 AC/DC
2 Accept
3 Aerosmith


CREATE TABLE "Track" (
	"TrackId" INTEGER NOT NULL, 
	"Name" NVARCHAR(200) NOT NULL, 
	"AlbumId" INTEGER, 
	"MediaTypeId" INTEGER NOT NULL, 
	"GenreId" INTEGER, 
	"Composer" NVARCHAR(220), 
	"Milliseconds" INTEGER NOT NULL, 
	"Bytes" INTEGER, 
	"UnitPrice" NUMERIC(10, 2) NOT NULL, 
	PRIMARY KEY ("TrackId"), 
	FOREIGN KEY("MediaTypeId") REFERE

'The top 3 best selling artists are Iron Maiden, U2, and Metallica.'

## Memory Use Case Example
In the following example, we provide a list of customized inputs to enable memory recording. AI will be able to remember the past conversation and use them as context for new answers.

In [1]:
import os
os.environ["OPENAI_API_KEY"] = "...input your openai api key..."

In [3]:
from langchain.agents.agent_toolkits import SQLDatabaseToolkit
from langchain.sql_database import SQLDatabase
from langchain.llms.openai import OpenAI

db = SQLDatabase.from_uri("sqlite:///../../../../../notebooks/Chinook.db")
llm = OpenAI(temperature=0, model_name="gpt-4")
toolkit = SQLDatabaseToolkit(db=db, llm=llm)


In [4]:
# We will re-use the default prefix, which you can find in ...agents.agent_toolkits.sql.prompt
SQL_SUFFIX = """Begin!
{chat_history}
Question: {input}
Thought: I should look at the tables in the database to see what I can query.
{agent_scratchpad}"""

input_variables=['chat_history', 'input', 'agent_scratchpad']

In [5]:
from langchain.memory import ConversationBufferMemory

memory = ConversationBufferMemory(memory_key="chat_history")

In [6]:
from langchain.agents import create_sql_agent

agent_executor = create_sql_agent(
    llm=llm,
    toolkit=toolkit,
    suffix=SQL_SUFFIX,
    input_variables=input_variables,
    verbose=True,
    agent_executor_kwargs={"memory": memory},
)

In [7]:
agent_executor.run("who's the top selling artist?")



[1m> Entering new AgentExecutor chain...[0m
[32;1m[1;3mAction: list_tables_sql_db
Action Input: [0m
Observation: [38;5;200m[1;3mPlaylist, Track, Album, InvoiceLine, Customer, MediaType, Invoice, PlaylistTrack, Artist, Genre, Employee[0m
Thought:[32;1m[1;3mI should check the schema of the Artist, Album, and Track tables to see if there's any information about sales or top selling artists.
Action: schema_sql_db
Action Input: Artist, Album, Track[0m
Observation: [33;1m[1;3m
CREATE TABLE "Artist" (
	"ArtistId" INTEGER NOT NULL, 
	"Name" NVARCHAR(120), 
	PRIMARY KEY ("ArtistId")
)

/*
3 rows from Artist table:
ArtistId	Name
1	AC/DC
2	Accept
3	Aerosmith
*/


CREATE TABLE "Album" (
	"AlbumId" INTEGER NOT NULL, 
	"Title" NVARCHAR(160) NOT NULL, 
	"ArtistId" INTEGER NOT NULL, 
	PRIMARY KEY ("AlbumId"), 
	FOREIGN KEY("ArtistId") REFERENCES "Artist" ("ArtistId")
)

/*
3 rows from Album table:
AlbumId	Title	ArtistId
1	For Those About To Rock We Salute You	1
2	Balls to the Wall	2
3	R

'Iron Maiden'

In [8]:
agent_executor.run("how many albums does he have?")



[1m> Entering new AgentExecutor chain...[0m
[32;1m[1;3mAction: list_tables_sql_db
Action Input: [0m
Observation: [38;5;200m[1;3mPlaylist, Track, Album, InvoiceLine, Customer, MediaType, Invoice, PlaylistTrack, Artist, Genre, Employee[0m
Thought:[32;1m[1;3mI should check the schema of the Album and Artist tables to see how I can find the number of albums for an artist.
Action: schema_sql_db
Action Input: Album, Artist[0m
Observation: [33;1m[1;3m
CREATE TABLE "Artist" (
	"ArtistId" INTEGER NOT NULL, 
	"Name" NVARCHAR(120), 
	PRIMARY KEY ("ArtistId")
)

/*
3 rows from Artist table:
ArtistId	Name
1	AC/DC
2	Accept
3	Aerosmith
*/


CREATE TABLE "Album" (
	"AlbumId" INTEGER NOT NULL, 
	"Title" NVARCHAR(160) NOT NULL, 
	"ArtistId" INTEGER NOT NULL, 
	PRIMARY KEY ("AlbumId"), 
	FOREIGN KEY("ArtistId") REFERENCES "Artist" ("ArtistId")
)

/*
3 rows from Album table:
AlbumId	Title	ArtistId
1	For Those About To Rock We Salute You	1
2	Balls to the Wall	2
3	Restless and Wild	2
*/[0m
T

'Iron Maiden has 21 albums.'