(prompt_db, db_id, db_path, limit_value=3, normalization=True)
| 130 | |
| 131 | |
| 132 | def extract_create_table_prompt(prompt_db, db_id, db_path, limit_value=3, normalization=True): |
| 133 | table_query = "SELECT * FROM sqlite_master WHERE type='table';" |
| 134 | tables = sqlite3.connect(db_path).cursor().execute(table_query).fetchall() |
| 135 | prompt = "" |
| 136 | table_names = [] |
| 137 | foreign_keys = [] |
| 138 | for table in tables: |
| 139 | table_name = table[1] |
| 140 | if table_name == 'sqlite_sequence': |
| 141 | continue |
| 142 | table_names.append(table_name) |
| 143 | if normalization: |
| 144 | table_name = table_name.lower() |
| 145 | create_table_statement = table[-1] |
| 146 | |
| 147 | # xzjin 解决 'CREATE TABLE rankings("ranking_date" DATE,"ranking" INT,"player_id" INT,"ranking_points" INT,"tours" INT,FOREIGN KEY(player_id) REFERENCES players(player_id))' |
| 148 | original_sql = create_table_statement.split("\n") |
| 149 | if len(original_sql) == 1: |
| 150 | original_sql = original_sql[0] |
| 151 | formatted_sql = re.sub(r'\((.*)\)', lambda match: f'(\n{match.group(1)}\n)', original_sql, flags=re.DOTALL) |
| 152 | # 在最后一个右括号前添加换行符 |
| 153 | formatted_sql = re.sub(r'(?<=\))$', '\n', formatted_sql) |
| 154 | formatted_sql = formatted_sql.replace(",", ",\n") |
| 155 | create_table_statement = formatted_sql |
| 156 | |
| 157 | # 外键表 |
| 158 | fk_create_table_statement = normalize_create_table(table_name, create_table_statement) |
| 159 | |
| 160 | # xzjin add 解决空行问题 |
| 161 | create_table_statement = create_table_statement.split("\n") |
| 162 | new_lines = [] |
| 163 | for line in create_table_statement: |
| 164 | if line.strip(): # 判断去除空白字符后是否还有内容 |
| 165 | new_lines.append(line) |
| 166 | create_table_statement = "\n".join(new_lines) |
| 167 | |
| 168 | table_info_query = f"PRAGMA table_info({table_name});" |
| 169 | top_k_row_query = f"SELECT * FROM {table_name} LIMIT {limit_value};" |
| 170 | headers = [x[1] for x in sqlite3.connect(db_path).cursor().execute(table_info_query).fetchall()] |
| 171 | |
| 172 | if "fk" in prompt_db.lower(): |
| 173 | if normalization.startswith("upper") or normalization.startswith("newupper"): |
| 174 | foreign_keys_one_table = get_foreign_keys(db_id, table_name, fk_create_table_statement, True) |
| 175 | else: |
| 176 | foreign_keys_one_table = get_foreign_keys(db_id, table_name, fk_create_table_statement, False) |
| 177 | foreign_keys.extend(foreign_keys_one_table) |
| 178 | |
| 179 | # xzjin add 控制空行和大小写 |
| 180 | if normalization == "upper": |
| 181 | create_table_statement = create_table_statement.split("\n") |
| 182 | new_lines = [] |
| 183 | for line in create_table_statement: |
| 184 | if "create table" in line.lower(): |
| 185 | line = line.upper() |
| 186 | line = re.sub(r'[`"\']', '', line) |
| 187 | new_lines.append(line) |
| 188 | create_table_statement = "\n".join(new_lines) |
| 189 |
no test coverage detected