diff --git a/tool/pputil/pputil.go b/tool/pputil/pputil.go index 85f501b5..42ce62b1 100644 --- a/tool/pputil/pputil.go +++ b/tool/pputil/pputil.go @@ -92,46 +92,136 @@ func fileFound(file, dir string) (string, bool) { // ----------------------------------------------------------------------------- -var include = []byte("#include") - // LoadIncludes loads the include directives from the specified header file. It // returns a sequence of Include objects. +// +// It parses the file in a single pass, correctly handling `#include` directives +// with spaces or tabs between `#` and `include`, and ignoring any `#include` +// occurrences found within comments or string/character literals. func LoadIncludes(headerFile string) (includes iter.Seq[Include], err error) { b, err := os.ReadFile(headerFile) if err != nil { return } includes = func(yield func(Include) bool) { - for { - pos := bytes.Index(b, include) - if pos < 0 { - break + scanIncludes(b, yield) + } + return +} + +var includeKw = []byte("include") + +// scanIncludes scans b for `#include` directives, yielding each one. It skips +// line comments, block comments, and string/character literals so that +// `#include` tokens appearing inside them are not mistaken for directives. +// +// A directive is recognized only when `#` is the first non-blank, non-comment +// token on a line, optionally followed by spaces or tabs, then `include`, then +// the header name in `"..."` or `<...>`. Comments are treated as whitespace, so +// a same-line block comment before the directive (e.g. `/* c */ #include `) +// does not suppress it. +func scanIncludes(b []byte, yield func(Include) bool) { + atLineStart := true // no non-blank character seen yet on the current line + for i := 0; i < len(b); { + c := b[i] + switch { + case c == '\n': + atLineStart = true + i++ + case c == ' ' || c == '\t' || c == '\r' || c == '\f' || c == '\v': + i++ // leading blanks keep us at line start + case c == '/' && i+1 < len(b) && b[i+1] == '/': + // line comment: skip to end of line + i += 2 + for i < len(b) && b[i] != '\n' { + i++ } - b = b[pos+len(include):] - b = bytes.TrimLeft(b, " \t") - if len(b) == 0 { - break + case c == '/' && i+1 < len(b) && b[i+1] == '*': + // block comment: skip to closing */. Comments count as + // whitespace, so atLineStart is left unchanged. + i += 2 + for i < len(b) && !(b[i] == '*' && i+1 < len(b) && b[i+1] == '/') { + i++ } - quote := b[0] - if quote != '"' { - if quote != '<' { - continue + i += 2 + case c == '"' || c == '\'': + i = skipLiteral(b, i, c) + atLineStart = false + case c == '#' && atLineStart: + if next, inc, ok := parseInclude(b, i); ok { + if !yield(inc) { + return } - quote = '>' - } - b = b[1:] - pos = bytes.IndexByte(b, quote) - if pos < 0 { - continue - } - fname := string(b[:pos]) - b = b[pos+1:] - if !yield(Include{Filename: fname, Quote: quote == '"'}) { - return + i = next + } else { + i++ } + atLineStart = false + default: + atLineStart = false + i++ } } - return +} + +// skipLiteral returns the index just past a string (") or character (') literal +// that starts at b[i] == quote, honoring backslash escapes. +func skipLiteral(b []byte, i int, quote byte) int { + i++ // opening quote + for i < len(b) { + switch b[i] { + case '\\': + i += 2 // skip the escaped character + case quote: + return i + 1 + case '\n': + return i // unterminated literal; stop at end of line + default: + i++ + } + } + return i +} + +// parseInclude tries to parse an `#include "..."` or `#include <...>` directive +// starting at b[i] == '#'. On success it returns the index just past the +// directive, the parsed Include, and true. +func parseInclude(b []byte, i int) (next int, inc Include, ok bool) { + j := i + 1 // past '#' + for j < len(b) && (b[j] == ' ' || b[j] == '\t') { + j++ + } + if !bytes.HasPrefix(b[j:], includeKw) { + return + } + j += len(includeKw) + // require a blank between `include` and the header name + start := j + for j < len(b) && (b[j] == ' ' || b[j] == '\t') { + j++ + } + if j == start || j >= len(b) { + return + } + open := b[j] + var closer byte + switch open { + case '"': + closer = '"' + case '<': + closer = '>' + default: + return + } + j++ + end := j + for end < len(b) && b[end] != closer && b[end] != '\n' { + end++ + } + if end >= len(b) || b[end] != closer { + return + } + return end + 1, Include{Filename: string(b[j:end]), Quote: open == '"'}, true } // ----------------------------------------------------------------------------- diff --git a/tool/pputil/pputil_test.go b/tool/pputil/pputil_test.go new file mode 100644 index 00000000..8aef2045 --- /dev/null +++ b/tool/pputil/pputil_test.go @@ -0,0 +1,218 @@ +/* + * Copyright (c) 2026 The XGo Authors (xgo.dev). All rights reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package pputil + +import ( + "os" + "path/filepath" + "reflect" + "testing" +) + +func scanAll(src string) []Include { + var got []Include + scanIncludes([]byte(src), func(inc Include) bool { + got = append(got, inc) + return true + }) + return got +} + +func TestScanIncludes(t *testing.T) { + tests := []struct { + name string + src string + want []Include + }{ + { + name: "basic quote and angle", + src: "#include \"foo.h\"\n#include \n", + want: []Include{ + {Filename: "foo.h", Quote: true}, + {Filename: "bar.h", Quote: false}, + }, + }, + { + name: "spaces and tabs after hash", + src: "# include \"a.h\"\n#\tinclude \n# \t include \"c.h\"\n", + want: []Include{ + {Filename: "a.h", Quote: true}, + {Filename: "b.h", Quote: false}, + {Filename: "c.h", Quote: true}, + }, + }, + { + name: "leading whitespace before hash", + src: " #include \"a.h\"\n\t#include \n", + want: []Include{ + {Filename: "a.h", Quote: true}, + {Filename: "b.h", Quote: false}, + }, + }, + { + name: "ignore line comment", + src: "// #include \"skip.h\"\n#include \"keep.h\"\nint x; // #include \n", + want: []Include{ + {Filename: "keep.h", Quote: true}, + }, + }, + { + name: "ignore block comment single line", + src: "/* #include \"skip.h\" */\n#include \"keep.h\"\n", + want: []Include{ + {Filename: "keep.h", Quote: true}, + }, + }, + { + name: "ignore block comment multi line", + src: "/*\n#include \"skip1.h\"\n#include \n*/\n#include \"keep.h\"\n", + want: []Include{ + {Filename: "keep.h", Quote: true}, + }, + }, + { + name: "same-line block comment before include", + src: "/* c */ #include \n/* lead */#include \"y.h\"\n", + want: []Include{ + {Filename: "x.h", Quote: false}, + {Filename: "y.h", Quote: true}, + }, + }, + { + name: "code before same-line block comment then include is ignored", + src: "int x; /* c */ #include \n#include \"keep.h\"\n", + want: []Include{ + {Filename: "keep.h", Quote: true}, + }, + }, + { + name: "ignore string literal", + src: "const char *s = \"#include \";\n#include \"real.h\"\n", + want: []Include{ + {Filename: "real.h", Quote: true}, + }, + }, + { + name: "ignore char literal and escapes", + src: "char c = '\"';\nchar *q = \"a\\\"b #include \";\n#include \"real.h\"\n", + want: []Include{ + {Filename: "real.h", Quote: true}, + }, + }, + { + name: "hash not at line start is ignored", + src: "int x = 1; #include \"skip.h\"\n#include \"keep.h\"\n", + want: []Include{ + {Filename: "keep.h", Quote: true}, + }, + }, + { + name: "no separator between include and name", + src: "#include\"a.h\"\n#include\n", + want: nil, + }, + { + name: "other directives ignored", + src: "#ifndef FOO\n#define FOO\n#include \"a.h\"\n#endif\n", + want: []Include{ + {Filename: "a.h", Quote: true}, + }, + }, + { + name: "unterminated header name ignored", + src: "#include \"a.h\n#include \n", + want: []Include{ + {Filename: "b.h", Quote: false}, + }, + }, + { + name: "no newline at eof", + src: "#include \"a.h\"", + want: []Include{ + {Filename: "a.h", Quote: true}, + }, + }, + { + name: "include keyword substring not matched", + src: "#includes \"a.h\"\n#include_next \n", + want: nil, + }, + { + name: "empty input", + src: "", + want: nil, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := scanAll(tt.src) + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("scanIncludes(%q) = %v, want %v", tt.src, got, tt.want) + } + }) + } +} + +func TestScanIncludesEarlyStop(t *testing.T) { + src := "#include \"a.h\"\n#include \"b.h\"\n#include \"c.h\"\n" + var got []Include + scanIncludes([]byte(src), func(inc Include) bool { + got = append(got, inc) + return len(got) < 2 // stop after two + }) + want := []Include{ + {Filename: "a.h", Quote: true}, + {Filename: "b.h", Quote: true}, + } + if !reflect.DeepEqual(got, want) { + t.Errorf("early stop = %v, want %v", got, want) + } +} + +func TestLoadIncludes(t *testing.T) { + dir := t.TempDir() + file := filepath.Join(dir, "test.h") + src := "// #include \"comment.h\"\n" + + "# include \"spaced.h\"\n" + + "const char *s = \"#include \";\n" + + "#include \n" + if err := os.WriteFile(file, []byte(src), 0644); err != nil { + t.Fatal(err) + } + includes, err := LoadIncludes(file) + if err != nil { + t.Fatal(err) + } + var got []Include + for inc := range includes { + got = append(got, inc) + } + want := []Include{ + {Filename: "spaced.h", Quote: true}, + {Filename: "sys.h", Quote: false}, + } + if !reflect.DeepEqual(got, want) { + t.Errorf("LoadIncludes = %v, want %v", got, want) + } +} + +func TestLoadIncludesError(t *testing.T) { + _, err := LoadIncludes(filepath.Join(t.TempDir(), "does-not-exist.h")) + if err == nil { + t.Error("expected error for missing file, got nil") + } +}