From 58cf80ec4bcf976960b4e5ca9eee830e4926c7eb Mon Sep 17 00:00:00 2001 From: Tomas Votruba Date: Fri, 25 Sep 2026 12:30:55 +0200 Subject: [PATCH 1/2] Skip string literal docblock when an unresolved value is passed, add --literals option --- CLAUDE.md | 3 ++- README.md | 11 +++++++-- internal/aggregate/aggregate.go | 29 ++++++++++++---------- internal/aggregate/aggregate_test.go | 19 +++++++++++++++ internal/apply/apply.go | 36 ++++++++++++++++++++-------- internal/apply/apply_test.go | 25 +++++++++++++++++-- internal/collect/collect.go | 11 ++++----- internal/collect/collect_test.go | 16 +++++++++---- main.go | 20 ++++++++++++---- 9 files changed, 128 insertions(+), 42 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index 2e488161..fa198c5d 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -43,10 +43,11 @@ a parent or interface (including vendor), unless private or a constructor. ## Commands ```bash -argtyper add-types [project-path] [--dry] +argtyper add-types [project-path] [--dry] [--literals] ``` `project-path` defaults to `.`. `--dry` prints the diff without writing. +`--literals` only adds string types with a `@param 'a'|'b'` literal docblock. Dev tasks via `Makefile`: diff --git a/README.md b/README.md index 6d6e3a80..d5b8da96 100644 --- a/README.md +++ b/README.md @@ -153,8 +153,9 @@ $this->compareScore(8, 'neq'); } ``` -Any non-literal string (a constant, `sprintf()`, interpolation) skips the -docblock, as does an existing `@param` for that parameter. +Any non-literal string (a constant, `sprintf()`, interpolation) or a value it +cannot resolve (a variable) skips the docblock, as does an existing `@param` +for that parameter. **Colored, informative output** - a live progress bar per phase, colored `--dry` diffs, and a summary of the added types grouped by category (scalar, object, @@ -194,6 +195,12 @@ argtyper add-types . --dry It prints the diff of the types that would be added and leaves every file untouched. +Only add the string literal docblocks, and no other types, with `--literals`: + +```bash +argtyper add-types . --literals +``` +
## How it works diff --git a/internal/aggregate/aggregate.go b/internal/aggregate/aggregate.go index 4f4d53a8..12bb60e5 100644 --- a/internal/aggregate/aggregate.go +++ b/internal/aggregate/aggregate.go @@ -25,9 +25,11 @@ const ( // group collects everything observed for one parameter position. type group struct { - types map[string]struct{} - literals map[string]struct{} - nonLiteralStrings bool + types map[string]struct{} + literals map[string]struct{} + // set when a string other than a plain literal, or an unresolved value, was + // passed - the literals then no longer cover every argument + nonLiterals bool } // Types holds resolved parameter types keyed for fast lookup during apply. @@ -98,10 +100,10 @@ func resolveGroups(groups map[string]*group) map[string]Resolved { } // literalValues returns the string literals passed into a plain string -// parameter, when every string argument was a literal and there are a few +// parameter, when every argument was a string literal or null and there are a few // distinct ones - an enum-like set worth a `'a'|'b'` doc type. func literalValues(members []string, group *group) []string { - if len(members) != 1 || members[0] != "string" || group.nonLiteralStrings { + if len(members) != 1 || members[0] != "string" || group.nonLiterals { return nil } if len(group.literals) < minLiterals || len(group.literals) > maxLiterals { @@ -119,15 +121,18 @@ func addRecord(groups map[string]*group, key string, record collect.Record) { if groups[key] == nil { groups[key] = &group{types: map[string]struct{}{}, literals: map[string]struct{}{}} } - groups[key].types[record.Type] = struct{}{} - if record.Type != "string" { - return - } - if record.IsLiteral { + switch { + case record.Type == "": + groups[key].nonLiterals = true + case record.IsLiteral: + groups[key].types[record.Type] = struct{}{} groups[key].literals[record.Literal] = struct{}{} - return + default: + groups[key].types[record.Type] = struct{}{} + if record.Type == "string" { + groups[key].nonLiterals = true + } } - groups[key].nonLiteralStrings = true } func methodKey(class, method string, position int) string { diff --git a/internal/aggregate/aggregate_test.go b/internal/aggregate/aggregate_test.go index 47104336..3d8a6c2e 100644 --- a/internal/aggregate/aggregate_test.go +++ b/internal/aggregate/aggregate_test.go @@ -52,6 +52,20 @@ func TestResolveMethod(t *testing.T) { wantNullable: true, wantFound: true, }, + { + name: "unresolved argument keeps the resolved type", + records: []collect.Record{ + {Class: "A", Name: "m", Type: "int"}, + {Class: "A", Name: "m"}, + }, + wantType: "int", + wantFound: true, + }, + { + name: "only unresolved resolves to nothing", + records: []collect.Record{{Class: "A", Name: "m"}}, + wantFound: false, + }, { name: "only null resolves to nothing", records: []collect.Record{{Class: "A", Name: "m", Type: "null"}}, @@ -127,6 +141,11 @@ func TestResolveLiterals(t *testing.T) { records: []collect.Record{literal("eq"), literal("neq"), {Class: "A", Name: "m", Type: "string"}}, want: "", }, + { + name: "unresolved argument drops the literals", + records: []collect.Record{literal("eq"), literal("neq"), {Class: "A", Name: "m"}}, + want: "", + }, { name: "union with another type drops the literals", records: []collect.Record{literal("eq"), literal("neq"), {Class: "A", Name: "m", Type: "int"}}, diff --git a/internal/apply/apply.go b/internal/apply/apply.go index c47845cc..7528db8b 100644 --- a/internal/apply/apply.go +++ b/internal/apply/apply.go @@ -19,14 +19,15 @@ import ( // "?string" or "\Foo"), and whether the file changed. On a parse error the // original source is returned unchanged. The symbols table resolves constant // and enum-case default values; the inheritance table decides whether typing a -// method would change an inherited signature. -func Source(src []byte, types aggregate.Types, table *symbols.Table, inheritance *inherit.Table) (string, []string, bool) { +// method would change an inherited signature. With literalsOnly, only string +// parameters that get a `'a'|'b'` literal doc type are changed. +func Source(src []byte, types aggregate.Types, table *symbols.Table, inheritance *inherit.Table, literalsOnly bool) (string, []string, bool) { root, err := phpast.Parse(src) if err != nil || root == nil { return string(src), nil, false } - applier := &applier{types: types, symbols: table, inheritance: inheritance, names: phpast.ResolveNames(root)} + applier := &applier{types: types, symbols: table, inheritance: inheritance, names: phpast.ResolveNames(root), literalsOnly: literalsOnly} applier.walk(root, nil) if len(applier.added) == 0 { @@ -37,11 +38,12 @@ func Source(src []byte, types aggregate.Types, table *symbols.Table, inheritance } type applier struct { - types aggregate.Types - symbols *symbols.Table - inheritance *inherit.Table - names map[ast.Vertex]string - added []string + types aggregate.Types + symbols *symbols.Table + inheritance *inherit.Table + names map[ast.Vertex]string + added []string + literalsOnly bool } func (a *applier) walk(node ast.Vertex, class *ast.StmtClass) { @@ -74,13 +76,20 @@ func (a *applier) applyFunction(function *ast.StmtFunction) { continue } if resolved, ok := a.types.Function(name, position); ok { + docType := literalDocType(param, resolved) + if a.literalsOnly && docType == "" { + continue + } added[phpast.VariableName(param.Var)] = a.setType(param, resolved) - if docType := literalDocType(param, resolved); docType != "" { + if docType != "" { docs[phpast.VariableName(param.Var)] = docType order = append(order, phpast.VariableName(param.Var)) } continue } + if a.literalsOnly { + continue + } if resolved, ok := a.defaultType(param, "", ""); ok { added[phpast.VariableName(param.Var)] = a.setType(param, resolved) } @@ -114,13 +123,20 @@ func (a *applier) applyMethod(method *ast.StmtClassMethod, class *ast.StmtClass) continue } if resolved, ok := a.types.Method(className, name, position); ok { + docType := literalDocType(param, resolved) + if a.literalsOnly && docType == "" { + continue + } added[phpast.VariableName(param.Var)] = a.setType(param, resolved) - if docType := literalDocType(param, resolved); docType != "" { + if docType != "" { docs[phpast.VariableName(param.Var)] = docType order = append(order, phpast.VariableName(param.Var)) } continue } + if a.literalsOnly { + continue + } if resolved, ok := a.defaultType(param, className, classFQCN); ok { added[phpast.VariableName(param.Var)] = a.setType(param, resolved) } diff --git a/internal/apply/apply_test.go b/internal/apply/apply_test.go index 0988a43b..a5c933ee 100644 --- a/internal/apply/apply_test.go +++ b/internal/apply/apply_test.go @@ -23,7 +23,7 @@ func run(target string, sources ...string) (string, int) { records = append(records, collect.FromSource([]byte(source), table)...) } types := aggregate.Resolve(records) - output, added, _ := apply.Source([]byte(target), types, table, inheritance) + output, added, _ := apply.Source([]byte(target), types, table, inheritance, false) return output, len(added) } @@ -324,6 +324,14 @@ func TestApply(t *testing.T) { }, want: "set($e['k'] ?? null); }\n}", - want: []collect.Record{{Class: "A", Name: "set", Position: 0, Type: "null"}}, + want: []collect.Record{ + {Class: "A", Name: "set", Position: 0}, + {Class: "A", Name: "set", Position: 0, Type: "null"}, + }, }, { name: "single quoted string literal keeps its value", diff --git a/main.go b/main.go index 6ea45574..1719c38f 100644 --- a/main.go +++ b/main.go @@ -16,12 +16,13 @@ import ( "github.com/rectorphp/argtyper/internal/symbols" ) -const usage = `Usage: argtyper add-types [project-path] [--dry] +const usage = `Usage: argtyper add-types [project-path] [--dry] [--literals] Find literal values passed into local method/function calls and add them as parameter type declarations. Defaults to the current directory. - --dry Print the diff of the types that would be added, without writing.` + --dry Print the diff of the types that would be added, without writing. + --literals Only add string types with a @param 'a'|'b' literal docblock.` func main() { if err := run(os.Args[1:]); err != nil { @@ -40,12 +41,17 @@ func run(args []string) error { } dry := false + literalsOnly := false projectPath := "." for _, arg := range args[1:] { if arg == "--dry" { dry = true continue } + if arg == "--literals" { + literalsOnly = true + continue + } projectPath = arg } @@ -90,7 +96,13 @@ func run(args []string) error { }); err != nil { return err } - fmt.Printf(" Found %d arg types\n\n", len(records)) + resolvedCount := 0 + for _, record := range records { + if record.Type != "" { + resolvedCount++ + } + } + fmt.Printf(" Found %d arg types\n\n", resolvedCount) types := aggregate.Resolve(records) @@ -102,7 +114,7 @@ func run(args []string) error { } var addedTypes []string if err := progressEach("applying", files, func(file string, src []byte) error { - output, added, changed := apply.Source(src, types, table, inheritance) + output, added, changed := apply.Source(src, types, table, inheritance, literalsOnly) if !changed { return nil } From b3a47f0fc6d6d27a8d62222fac6f02551d430754 Mon Sep 17 00:00:00 2001 From: Tomas Votruba Date: Fri, 25 Sep 2026 17:30:07 +0200 Subject: [PATCH 2/2] Add --objects option to add only object types (#40) --- CLAUDE.md | 4 ++- README.md | 9 ++++++ internal/apply/apply.go | 62 ++++++++++++++++++++++++------------ internal/apply/apply_test.go | 33 +++++++++++++++++-- main.go | 17 +++++++--- 5 files changed, 96 insertions(+), 29 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index fa198c5d..4a5b2af9 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -43,11 +43,13 @@ a parent or interface (including vendor), unless private or a constructor. ## Commands ```bash -argtyper add-types [project-path] [--dry] [--literals] +argtyper add-types [project-path] [--dry] [--literals] [--objects] ``` `project-path` defaults to `.`. `--dry` prints the diff without writing. `--literals` only adds string types with a `@param 'a'|'b'` literal docblock. +`--objects` only adds object types, including unions made only of objects. +Both can be combined. Dev tasks via `Makefile`: diff --git a/README.md b/README.md index d5b8da96..c5babd3d 100644 --- a/README.md +++ b/README.md @@ -201,6 +201,15 @@ Only add the string literal docblocks, and no other types, with `--literals`: argtyper add-types . --literals ``` +Only add object types with `--objects` - a single class, a nullable one, or a +union made only of classes (`\Foo|\Bar`). A union with a scalar, like +`int|\Money`, is skipped. Both options can be combined: + +```bash +argtyper add-types . --objects +argtyper add-types . --literals --objects +``` +
## How it works diff --git a/internal/apply/apply.go b/internal/apply/apply.go index 7528db8b..925c57ea 100644 --- a/internal/apply/apply.go +++ b/internal/apply/apply.go @@ -19,15 +19,15 @@ import ( // "?string" or "\Foo"), and whether the file changed. On a parse error the // original source is returned unchanged. The symbols table resolves constant // and enum-case default values; the inheritance table decides whether typing a -// method would change an inherited signature. With literalsOnly, only string -// parameters that get a `'a'|'b'` literal doc type are changed. -func Source(src []byte, types aggregate.Types, table *symbols.Table, inheritance *inherit.Table, literalsOnly bool) (string, []string, bool) { +// method would change an inherited signature. The options limit which types +// are added. +func Source(src []byte, types aggregate.Types, table *symbols.Table, inheritance *inherit.Table, options Options) (string, []string, bool) { root, err := phpast.Parse(src) if err != nil || root == nil { return string(src), nil, false } - applier := &applier{types: types, symbols: table, inheritance: inheritance, names: phpast.ResolveNames(root), literalsOnly: literalsOnly} + applier := &applier{types: types, symbols: table, inheritance: inheritance, names: phpast.ResolveNames(root), options: options} applier.walk(root, nil) if len(applier.added) == 0 { @@ -37,13 +37,41 @@ func Source(src []byte, types aggregate.Types, table *symbols.Table, inheritance return phpast.Print(root), applier.added, true } +// Options limit the added types to some kinds. With none set, every type is +// added; with one or more set, only the chosen kinds are. +type Options struct { + Literals bool // string types that get a `'a'|'b'` literal doc type + Objects bool // object types, including unions made only of objects +} + +// allows reports whether a resolved type passes the options; docType is its +// literal doc type, if any. +func (o Options) allows(resolved aggregate.Resolved, docType string) bool { + if !o.Literals && !o.Objects { + return true + } + if o.Literals && docType != "" { + return true + } + return o.Objects && onlyObjects(resolved) +} + +func onlyObjects(resolved aggregate.Resolved) bool { + for _, member := range resolved.Types { + if !strings.HasPrefix(member, "object:") { + return false + } + } + return len(resolved.Types) > 0 +} + type applier struct { - types aggregate.Types - symbols *symbols.Table - inheritance *inherit.Table - names map[ast.Vertex]string - added []string - literalsOnly bool + types aggregate.Types + symbols *symbols.Table + inheritance *inherit.Table + names map[ast.Vertex]string + added []string + options Options } func (a *applier) walk(node ast.Vertex, class *ast.StmtClass) { @@ -77,7 +105,7 @@ func (a *applier) applyFunction(function *ast.StmtFunction) { } if resolved, ok := a.types.Function(name, position); ok { docType := literalDocType(param, resolved) - if a.literalsOnly && docType == "" { + if !a.options.allows(resolved, docType) { continue } added[phpast.VariableName(param.Var)] = a.setType(param, resolved) @@ -87,10 +115,7 @@ func (a *applier) applyFunction(function *ast.StmtFunction) { } continue } - if a.literalsOnly { - continue - } - if resolved, ok := a.defaultType(param, "", ""); ok { + if resolved, ok := a.defaultType(param, "", ""); ok && a.options.allows(resolved, "") { added[phpast.VariableName(param.Var)] = a.setType(param, resolved) } } @@ -124,7 +149,7 @@ func (a *applier) applyMethod(method *ast.StmtClassMethod, class *ast.StmtClass) } if resolved, ok := a.types.Method(className, name, position); ok { docType := literalDocType(param, resolved) - if a.literalsOnly && docType == "" { + if !a.options.allows(resolved, docType) { continue } added[phpast.VariableName(param.Var)] = a.setType(param, resolved) @@ -134,10 +159,7 @@ func (a *applier) applyMethod(method *ast.StmtClassMethod, class *ast.StmtClass) } continue } - if a.literalsOnly { - continue - } - if resolved, ok := a.defaultType(param, className, classFQCN); ok { + if resolved, ok := a.defaultType(param, className, classFQCN); ok && a.options.allows(resolved, "") { added[phpast.VariableName(param.Var)] = a.setType(param, resolved) } } diff --git a/internal/apply/apply_test.go b/internal/apply/apply_test.go index a5c933ee..095b2a38 100644 --- a/internal/apply/apply_test.go +++ b/internal/apply/apply_test.go @@ -23,7 +23,7 @@ func run(target string, sources ...string) (string, int) { records = append(records, collect.FromSource([]byte(source), table)...) } types := aggregate.Resolve(records) - output, added, _ := apply.Source([]byte(target), types, table, inheritance, false) + output, added, _ := apply.Source([]byte(target), types, table, inheritance, apply.Options{}) return output, len(added) } @@ -354,7 +354,7 @@ func TestApply(t *testing.T) { func TestNoTypesReturnsUnchanged(t *testing.T) { src := "