-
Notifications
You must be signed in to change notification settings - Fork 64
Expand file tree
/
Copy pathresponse.py
More file actions
153 lines (117 loc) · 4.39 KB
/
Copy pathresponse.py
File metadata and controls
153 lines (117 loc) · 4.39 KB
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
import json
import re
import black
def wrap_code(code: str, lang="python") -> str:
return f"```{lang}\n{code}\n```"
def is_valid_python_script(script):
try:
compile(script, "<string>", "exec")
return True
except SyntaxError:
return False
def extract_jsons(text):
json_objects = []
matches = re.findall(r"\{.*?\}", text, re.DOTALL)
for match in matches:
try:
json_obj = json.loads(match)
json_objects.append(json_obj)
except json.JSONDecodeError:
pass
if len(json_objects) == 0 and not text.endswith("}"):
json_objects = extract_jsons(text + "}")
if len(json_objects) > 0:
return json_objects
return json_objects
def trim_long_string(string, threshold=5100, k=2500):
"""Truncate to first-k + last-k; rescue `Final ... Validation ... : <num>` lines into the middle."""
if len(string) <= threshold:
return string
strict = re.compile(r'^Final\s+Validation\s+\w+\s*[:=]\s*[-+]?\d', re.IGNORECASE | re.MULTILINE)
key_lines = [l for l in string.split('\n') if strict.search(l)]
if not key_lines:
loose = re.compile(r'Final\s+[\w\s]*?Validation\s+[\w\s]*?[:=]\s*[-+]?\d', re.IGNORECASE)
key_lines = [l for l in string.split('\n') if loose.search(l)]
first_k_chars = string[:k]
last_k_chars = string[-k:]
truncated_len = len(string) - 2 * k
if key_lines:
key_block = '\n'.join(key_lines[-3:])
return (
f"{first_k_chars}\n"
f" ... [{truncated_len} characters truncated] ... \n"
f"{key_block}\n"
f" ... [output continues] ... \n"
f"{last_k_chars}"
)
return f"{first_k_chars}\n ... [{truncated_len} characters truncated] ... \n{last_k_chars}"
def extract_code(text):
parsed_codes = []
matches = re.findall(r"```(python)?\n*(.*?)\n*```", text, re.DOTALL)
for match in matches:
code_block = match[1]
parsed_codes.append(code_block)
if len(parsed_codes) == 0:
matches = re.findall(r"^(```(python)?)?\n?(.*?)\n?(```)?$", text, re.DOTALL)
if matches:
code_block = matches[0][2]
parsed_codes.append(code_block)
valid_code_blocks = [
format_code(c) for c in parsed_codes if is_valid_python_script(c)
]
return format_code("\n\n".join(valid_code_blocks))
def extract_text_up_to_code(s):
if "```" not in s:
return ""
return s[: s.find("```")].strip()
def extract_plan_from_diff_response(text: str) -> str:
if not text:
return ""
stop_tokens = [
"<<<<<<< SEARCH",
"< SEARCH",
">>>>>>> REPLACE",
"=======",
"```",
]
def cut_at_stop(s: str) -> str:
indices = [s.find(token) for token in stop_tokens if s.find(token) != -1]
if indices:
return s[: min(indices)]
return s
if "Fixed Code Plan:" in text:
candidate = text.split("Fixed Code Plan:", 1)[1]
return cut_at_stop(candidate).strip()
if "Plan:" in text:
candidate = text.split("Plan:", 1)[1]
return cut_at_stop(candidate).strip()
return cut_at_stop(text).strip()
def extract_review(text):
parsed_codes = []
matches = re.findall(r"```(json)?\n*(.*?)\n*```", text, re.DOTALL)
for match in matches:
code_block = match[1]
parsed_codes.append(code_block)
if len(parsed_codes) == 0:
matches = re.findall(r"^(```(json)?)?\n?(.*?)\n?(```)?$", text, re.DOTALL)
if matches:
code_block = matches[0][2]
parsed_codes.append(code_block)
if len(parsed_codes) == 0 or not parsed_codes[0].strip():
json_objects = extract_jsons(text)
if len(json_objects) > 0:
return json_objects[0]
raise ValueError(f"No JSON found in text")
try:
review = json.loads(parsed_codes[0].strip())
return review
except json.JSONDecodeError:
json_objects = extract_jsons(text)
if len(json_objects) > 0:
return json_objects[0]
raise
def format_code(code) -> str:
try:
return black.format_str(code, mode=black.FileMode())
except black.parsing.InvalidInput: # type: ignore
return code