MCPcopy Create free account
hub / github.com/FlyingFeather/DEA-SQL / extract_create_table_prompt

Function extract_create_table_prompt

data_preprocess.py:132–288  ·  view source on GitHub ↗
(prompt_db, db_id, db_path, limit_value=3, normalization=True)

Source from the content-addressed store, hash-verified

130
131
132def 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'[`"\']', '&#x27;, line)
187 new_lines.append(line)
188 create_table_statement = "\n".join(new_lines)
189

Callers 1

Calls 4

normalize_create_tableFunction · 0.85
get_foreign_keysFunction · 0.85
is_numberFunction · 0.85
executeMethod · 0.45

Tested by

no test coverage detected