File size: 5,156 Bytes
4fcd019
2020fa0
 
 
 
 
 
 
 
 
4fcd019
 
 
2020fa0
4fcd019
2020fa0
4fcd019
2020fa0
 
 
4fcd019
 
 
2020fa0
 
 
 
4fcd019
 
 
2020fa0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4fcd019
2020fa0
4fcd019
 
 
2020fa0
 
 
 
 
 
4fcd019
 
 
2020fa0
 
 
4fcd019
2020fa0
4fcd019
2020fa0
4fcd019
 
 
2020fa0
 
 
 
 
7b5809f
 
 
2020fa0
 
 
 
 
 
7b5809f
 
 
2020fa0
 
 
 
 
0212283
 
 
64cb3d4
 
 
0212283
 
 
 
 
 
 
 
 
 
 
64cb3d4
 
4fcd019
 
 
2020fa0
4fcd019
2020fa0
 
4fcd019
 
 
 
 
 
2020fa0
4fcd019
2020fa0
4fcd019
 
 
 
 
35212b7
2020fa0
35212b7
2020fa0
35212b7
 
 
2020fa0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
"""
Database layer for TalkToDoc.
Local development and testing use SQLite (no setup needed). When deployed,
DATABASE_URL is set by Render and the app uses Postgres instead, since
Render's free tier wipes local files like a SQLite database on every
restart.

Queries are written once using SQLite-style '?' placeholders and adapted
automatically for Postgres, so there's a single query per function rather
than two versions of everything.
"""

import os
from pathlib import Path

import sqlite3

DB_PATH = Path(__file__).parent / "talktodoc.db"
SCHEMA_PATH = Path(__file__).parent / "schema.sql"
SCHEMA_PATH_POSTGRES = Path(__file__).parent / "schema_postgres.sql"

DATABASE_URL = os.environ.get("DATABASE_URL")

if DATABASE_URL:
    import psycopg2
    import psycopg2.extras


def get_connection():
    if DATABASE_URL:
        connection = psycopg2.connect(DATABASE_URL, cursor_factory=psycopg2.extras.RealDictCursor)
        # psycopg2 does not auto-commit by default. Without this, every
        # INSERT and UPDATE is silently rolled back when the connection
        # closes, so nothing is ever actually written to Postgres.
        # SQLite's context manager commits on clean exit, so this brings
        # the Postgres path in line with SQLite's behaviour.
        connection.autocommit = True
        return connection
    connection = sqlite3.connect(DB_PATH)
    connection.row_factory = sqlite3.Row
    connection.execute("PRAGMA foreign_keys = ON")
    return connection


def _run(connection, query, params=()):
    if DATABASE_URL:
        query = query.replace("?", "%s")
    cursor = connection.cursor()
    cursor.execute(query, params)
    return cursor


def _insert_and_get_id(connection, query, params):
    if DATABASE_URL:
        cursor = _run(connection, query + " RETURNING id", params)
        return cursor.fetchone()["id"]
    cursor = _run(connection, query, params)
    return cursor.lastrowid


def init_db():
    schema_path = SCHEMA_PATH_POSTGRES if DATABASE_URL else SCHEMA_PATH
    schema_sql = schema_path.read_text()
    with get_connection() as connection:
        if DATABASE_URL:
            connection.cursor().execute(schema_sql)
        else:
            connection.executescript(schema_sql)


def add_user(name, preferred_language, role):
    with get_connection() as connection:
        return _insert_and_get_id(
            connection,
            "INSERT INTO app_user (name, preferred_language, role) VALUES (?, ?, ?)",
            (name, preferred_language, role),
        )


def add_session(user_id, start_time):
    with get_connection() as connection:
        return _insert_and_get_id(
            connection,
            "INSERT INTO session (user_id, start_time) VALUES (?, ?)",
            (user_id, start_time),
        )


def end_session(session_id, end_time):
    with get_connection() as connection:
        _run(connection, "UPDATE session SET end_time = ? WHERE id = ?", (end_time, session_id))


def add_interaction(user_id, input_text, detected_language, translated_text,
                    nlu_summary, timestamp):
    sql = (
        "INSERT INTO interaction "
        "(user_id, input_text, detected_language, translated_text, nlu_summary, timestamp) "
        "VALUES (?, ?, ?, ?, ?, ?)"
    )

    params = (
        user_id,
        input_text,
        detected_language,
        translated_text,
        nlu_summary,
        timestamp
    )

    with get_connection() as connection:
        return _insert_and_get_id(connection, sql, params)


def get_interactions_for_user(user_id):
    with get_connection() as connection:
        cursor = _run(
            connection,
            "SELECT * FROM interaction WHERE user_id = ? ORDER BY timestamp",
            (user_id,),
        )
        return [dict(row) for row in cursor.fetchall()]


def get_pending_interactions():
    with get_connection() as connection:
        cursor = _run(
            connection,
            "SELECT * FROM interaction WHERE provider_response IS NULL ORDER BY timestamp",
        )
        return [dict(row) for row in cursor.fetchall()]


def get_completed_interactions():
    with get_connection() as connection:
        cursor = _run(
            connection,
            "SELECT * FROM interaction WHERE provider_response IS NOT NULL ORDER BY timestamp DESC",
        )
        return [dict(row) for row in cursor.fetchall()]


def get_interaction(interaction_id):
    with get_connection() as connection:
        cursor = _run(connection, "SELECT * FROM interaction WHERE id = ?", (interaction_id,))
        row = cursor.fetchone()
        return dict(row) if row else None


def update_interaction_response(interaction_id, provider_response, translated_response):
    with get_connection() as connection:
        _run(
            connection,
            "UPDATE interaction SET provider_response = ?, translated_response = ? WHERE id = ?",
            (provider_response, translated_response, interaction_id),
        )


if __name__ == "__main__":
    init_db()
    if DATABASE_URL:
        print("Database initialized (Postgres)")
    else:
        print("Database created at:", DB_PATH)