diff --git a/expander.py b/expander.py index e4b87f3..a5fa077 100755 --- a/expander.py +++ b/expander.py @@ -27,6 +27,30 @@ def is_ignored_line(self, line) -> bool: return True return False + def in_block_comment_after(self, line: str, in_comment: bool) -> bool: + # returns whether a /* */ comment is still open at the end of line + i = 0 + quote = None # type: Optional[str] + while i < len(line): + if in_comment: + if line.startswith('*/', i): + in_comment = False + i += 1 + elif quote: + if line[i] == '\\': + i += 1 + elif line[i] == quote: + quote = None + elif line.startswith('//', i): + break + elif line.startswith('/*', i): + in_comment = True + i += 1 + elif line[i] in '"\'': + quote = line[i] + i += 1 + return in_comment + def __init__(self, lib_paths: List[Path]): self.lib_paths = lib_paths @@ -67,9 +91,11 @@ def expand(self, source: str, origname) -> str: self.included = set() result = [] # type: List[str] linenum = 0 + in_comment = False for line in source.splitlines(): linenum += 1 - m = self.atcoder_include.match(line) + m = None if in_comment else self.atcoder_include.match(line) + in_comment = self.in_block_comment_after(line, in_comment) if m: acl_path = self.find_acl(m.group(1)) result.extend(self.expand_acl(acl_path)) diff --git a/test/expander/comment_out_multiline.cpp b/test/expander/comment_out_multiline.cpp new file mode 100644 index 0000000..2b21232 --- /dev/null +++ b/test/expander/comment_out_multiline.cpp @@ -0,0 +1,12 @@ +/* +#include +*/ +/* comment */ /* +#include +*/ +const char *s = "/*"; +#include + +int main() { + atcoder::dsu uf(10); +} diff --git a/test/test_expander.py b/test/test_expander.py index 28518ea..e278fe7 100755 --- a/test/test_expander.py +++ b/test/test_expander.py @@ -42,6 +42,10 @@ def test_comment_out(self): self.compile_test(Path('test/expander/comment_out.cpp'), expander_args=['--lib', str(Path.cwd().resolve())]) + def test_comment_out_multiline(self): + self.compile_test(Path('test/expander/comment_out_multiline.cpp'), + expander_args=['--lib', str(Path.cwd().resolve())]) + def test_env_value(self): env = environ.copy() env['CPLUS_INCLUDE_PATH'] = str(Path.cwd().resolve())