From 97c0859609fac047054f287acc0e135a0efd871b Mon Sep 17 00:00:00 2001 From: Tomas Votruba Date: Fri, 25 Sep 2026 11:59:58 +0200 Subject: [PATCH 1/2] Add @param string literal union when only literals are passed --- README.md | 22 ++++++ internal/aggregate/aggregate.go | 63 ++++++++++++--- internal/aggregate/aggregate_test.go | 55 +++++++++++++ internal/apply/apply.go | 40 ++++++++++ internal/apply/apply_test.go | 65 ++++++++++++++++ internal/collect/collect.go | 7 ++ internal/collect/collect_test.go | 20 ++++- internal/phpast/phpast.go | 112 +++++++++++++++++++++++++++ 8 files changed, 371 insertions(+), 13 deletions(-) diff --git a/README.md b/README.md index 61b8d35d..6d6e3a80 100644 --- a/README.md +++ b/README.md @@ -134,6 +134,28 @@ untyped. is removed, including `@param mixed`. Matching ignores union member order and FQN vs short name, so `@param Foo|Bar` and `@param \App\Bar|\App\Foo` both drop. +**String literal docblocks** - when a parameter only ever receives a few plain +string literals (2 to 10 distinct ones), it gets `string` plus a `@param` with +the exact values: + +```php +$this->compareScore(7, 'eq'); +$this->compareScore(8, 'neq'); +``` + +```diff ++/** ++ * @param 'eq'|'neq' $operator ++ */ +-public function compareScore(int $score, $operator) ++public function compareScore(int $score, string $operator) + { + } +``` + +Any non-literal string (a constant, `sprintf()`, interpolation) 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, array, union). Colors respect `NO_COLOR` and disable on non-TTY output. diff --git a/internal/aggregate/aggregate.go b/internal/aggregate/aggregate.go index 060210fe..4f4d53a8 100644 --- a/internal/aggregate/aggregate.go +++ b/internal/aggregate/aggregate.go @@ -15,6 +15,19 @@ import ( type Resolved struct { Types []string // one or more source type keywords or "object:Fqcn", never "null" Nullable bool + Literals []string // sorted string literal values, set only for a plain string type, see literalValues +} + +const ( + minLiterals = 2 + maxLiterals = 10 +) + +// group collects everything observed for one parameter position. +type group struct { + types map[string]struct{} + literals map[string]struct{} + nonLiteralStrings bool } // Types holds resolved parameter types keyed for fast lookup during apply. @@ -38,17 +51,17 @@ func (t Types) Function(name string, position int) (Resolved, bool) { // Resolve groups records per parameter into a resolved type: the observed types // as a union, with null captured as nullability rather than a member. func Resolve(records []collect.Record) Types { - methodTypes := map[string]map[string]struct{}{} - functionTypes := map[string]map[string]struct{}{} + methodTypes := map[string]*group{} + functionTypes := map[string]*group{} for _, record := range records { if record.IsFunction { key := functionKey(record.Name, record.Position) - addType(functionTypes, key, record.Type) + addRecord(functionTypes, key, record) continue } key := methodKey(record.Class, record.Name, record.Position) - addType(methodTypes, key, record.Type) + addRecord(methodTypes, key, record) } return Types{ @@ -57,12 +70,12 @@ func Resolve(records []collect.Record) Types { } } -func resolveGroups(groups map[string]map[string]struct{}) map[string]Resolved { +func resolveGroups(groups map[string]*group) map[string]Resolved { resolved := map[string]Resolved{} - for key, typeSet := range groups { - types := make([]string, 0, len(typeSet)) - for typeName := range typeSet { + for key, group := range groups { + types := make([]string, 0, len(group.types)) + for typeName := range group.types { types = append(types, typeName) } sort.Strings(types) @@ -78,17 +91,43 @@ func resolveGroups(groups map[string]map[string]struct{}) map[string]Resolved { if len(members) == 0 { continue } - resolved[key] = Resolved{Types: members, Nullable: nullable} + resolved[key] = Resolved{Types: members, Nullable: nullable, Literals: literalValues(members, group)} } return resolved } -func addType(groups map[string]map[string]struct{}, key, typeName string) { +// literalValues returns the string literals passed into a plain string +// parameter, when every string argument was a literal 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 { + return nil + } + if len(group.literals) < minLiterals || len(group.literals) > maxLiterals { + return nil + } + literals := make([]string, 0, len(group.literals)) + for literal := range group.literals { + literals = append(literals, literal) + } + sort.Strings(literals) + return literals +} + +func addRecord(groups map[string]*group, key string, record collect.Record) { if groups[key] == nil { - groups[key] = map[string]struct{}{} + 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 { + groups[key].literals[record.Literal] = struct{}{} + return } - groups[key][typeName] = struct{}{} + 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 ba9dabe4..47104336 100644 --- a/internal/aggregate/aggregate_test.go +++ b/internal/aggregate/aggregate_test.go @@ -88,3 +88,58 @@ func TestResolveFunctionAndPositionIsolation(t *testing.T) { t.Errorf("position 1: got %+v ok=%v", second, ok) } } + +func TestResolveLiterals(t *testing.T) { + literal := func(value string) collect.Record { + return collect.Record{Class: "A", Name: "m", Type: "string", IsLiteral: true, Literal: value} + } + + tests := []struct { + name string + records []collect.Record + want string // literals joined by "|" + }{ + { + name: "distinct literals sorted and deduplicated", + records: []collect.Record{literal("neq"), literal("eq"), literal("neq")}, + want: "eq|neq", + }, + { + name: "literals with null", + records: []collect.Record{literal("eq"), literal("neq"), {Class: "A", Name: "m", Type: "null"}}, + want: "eq|neq", + }, + { + name: "single literal is skipped", + records: []collect.Record{literal("eq"), literal("eq")}, + want: "", + }, + { + name: "more than ten literals are skipped", + records: []collect.Record{ + literal("a"), literal("b"), literal("c"), literal("d"), literal("e"), literal("f"), + literal("g"), literal("h"), literal("i"), literal("j"), literal("k"), + }, + want: "", + }, + { + name: "non-literal string drops the literals", + records: []collect.Record{literal("eq"), literal("neq"), {Class: "A", Name: "m", Type: "string"}}, + want: "", + }, + { + name: "union with another type drops the literals", + records: []collect.Record{literal("eq"), literal("neq"), {Class: "A", Name: "m", Type: "int"}}, + want: "", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + resolved, _ := aggregate.Resolve(test.records).Method("A", "m", 0) + if got := strings.Join(resolved.Literals, "|"); got != test.want { + t.Errorf("got %q want %q", got, test.want) + } + }) + } +} diff --git a/internal/apply/apply.go b/internal/apply/apply.go index 1bf7615e..c47845cc 100644 --- a/internal/apply/apply.go +++ b/internal/apply/apply.go @@ -3,6 +3,7 @@ package apply import ( + "slices" "strings" "github.com/rectorphp/argtyper/internal/aggregate" @@ -65,6 +66,8 @@ func (a *applier) walk(node ast.Vertex, class *ast.StmtClass) { func (a *applier) applyFunction(function *ast.StmtFunction) { name := phpast.ShortName(function.Name) added := map[string]string{} + docs := map[string]string{} + var order []string for position, paramNode := range function.Params { param, ok := paramNode.(*ast.Parameter) if !ok || !typeable(param) { @@ -72,6 +75,10 @@ func (a *applier) applyFunction(function *ast.StmtFunction) { } if resolved, ok := a.types.Function(name, position); ok { added[phpast.VariableName(param.Var)] = a.setType(param, resolved) + if docType := literalDocType(param, resolved); docType != "" { + docs[phpast.VariableName(param.Var)] = docType + order = append(order, phpast.VariableName(param.Var)) + } continue } if resolved, ok := a.defaultType(param, "", ""); ok { @@ -79,6 +86,7 @@ func (a *applier) applyFunction(function *ast.StmtFunction) { } } phpast.StripRedundantDocParams(function, added) + phpast.AddDocParams(function, docs, order) } func (a *applier) applyMethod(method *ast.StmtClassMethod, class *ast.StmtClass) { @@ -98,6 +106,8 @@ func (a *applier) applyMethod(method *ast.StmtClassMethod, class *ast.StmtClass) } added := map[string]string{} + docs := map[string]string{} + var order []string for position, paramNode := range method.Params { param, ok := paramNode.(*ast.Parameter) if !ok || !typeable(param) { @@ -105,6 +115,10 @@ func (a *applier) applyMethod(method *ast.StmtClassMethod, class *ast.StmtClass) } if resolved, ok := a.types.Method(className, name, position); ok { added[phpast.VariableName(param.Var)] = a.setType(param, resolved) + if docType := literalDocType(param, resolved); docType != "" { + docs[phpast.VariableName(param.Var)] = docType + order = append(order, phpast.VariableName(param.Var)) + } continue } if resolved, ok := a.defaultType(param, className, classFQCN); ok { @@ -112,6 +126,7 @@ func (a *applier) applyMethod(method *ast.StmtClassMethod, class *ast.StmtClass) } } phpast.StripRedundantDocParams(method, added) + phpast.AddDocParams(method, docs, order) } // defaultType infers a parameter type from its literal default value, so @@ -134,6 +149,31 @@ func (a *applier) defaultType(param *ast.Parameter, enclosing, enclosingFQCN str return aggregate.Resolved{Types: []string{typeName}}, true } +// literalDocType returns a `'a'|'b'` doc type for a parameter that only ever +// receives a few string literals. Empty when there are none, or when a default +// value falls outside them. +func literalDocType(param *ast.Parameter, resolved aggregate.Resolved) string { + if len(resolved.Literals) == 0 { + return "" + } + if param.DefaultValue != nil && !hasNullDefault(param) { + value, ok := phpast.StringLiteral(param.DefaultValue) + if !ok || !slices.Contains(resolved.Literals, value) { + return "" + } + } + + members := make([]string, len(resolved.Literals)) + for i, literal := range resolved.Literals { + members[i] = "'" + literal + "'" + } + docType := strings.Join(members, "|") + if resolved.Nullable || hasNullDefault(param) { + docType += "|null" + } + return docType +} + // objectFQCN qualifies the class of a `new X()` or `X::CASE` default value, // mapping self/static to the enclosing class. Empty when it cannot be resolved. func (a *applier) objectFQCN(expr ast.Vertex, enclosingFQCN string) string { diff --git a/internal/apply/apply_test.go b/internal/apply/apply_test.go index a303faee..0988a43b 100644 --- a/internal/apply/apply_test.go +++ b/internal/apply/apply_test.go @@ -267,6 +267,71 @@ func TestApply(t *testing.T) { target: "save($lead); }\n}", want: "save($lead); }\n}", }, + { + name: "adds string literals doc when only literals are passed", + target: "compare(7, 'eq'); $this->compare(8, 'neq'); $this->compare(6, 'gt'); }\n}", + want: "compare(7, 'eq'); $this->compare(8, 'neq'); $this->compare(6, 'gt'); }\n}", + }, + { + name: "adds string literals doc to an existing doc comment", + target: "set('a'); $this->set('b'); }\n}", + want: "set('a'); $this->set('b'); }\n}", + }, + { + name: "expands a single-line doc comment for string literals", + target: "set('a'); $this->set('b'); }\n}", + want: "set('a'); $this->set('b'); }\n}", + }, + { + name: "replaces a redundant string doc with string literals", + target: "set('a'); $this->set('b'); }\n}", + want: "set('a'); $this->set('b'); }\n}", + }, + { + name: "keeps an existing param doc with description over string literals", + target: "set('a'); $this->set('b'); }\n}", + want: "set('a'); $this->set('b'); }\n}", + }, + { + name: "adds nullable string literals doc", + target: "set($e['k'] ?? null); }\n}", want: []collect.Record{{Class: "A", Name: "set", Position: 0, Type: "null"}}, }, + { + name: "single quoted string literal keeps its value", + src: " $` lines to a function or method's doc +// comment, creating the doc comment when there is none. Parameters that already +// have a @param line are left alone. params maps parameter name to doc type and +// order lists the names in parameter order. +func AddDocParams(node ast.Vertex, params map[string]string, order []string) { + leading := leadingToken(node) + if leading == nil || len(params) == 0 { + return + } + + index := -1 + for i, free := range leading.FreeFloating { + if free.ID == token.T_DOC_COMMENT { + index = i + break + } + } + + doc := "" + if index >= 0 { + doc = string(leading.FreeFloating[index].Value) + } + + var tags []string + for _, name := range order { + docType, ok := params[name] + if !ok || hasDocParam(doc, name) { + continue + } + tags = append(tags, "@param "+docType+" $"+name) + } + if len(tags) == 0 { + return + } + + indent := indentation(leading.FreeFloating) + if index >= 0 { + leading.FreeFloating[index].Value = []byte(appendDocTags(doc, tags, indent)) + return + } + + newDoc := appendDocTags("/**\n"+indent+" */", tags, indent) + leading.FreeFloating = append(leading.FreeFloating, + &token.Token{ID: token.T_DOC_COMMENT, Value: []byte(newDoc)}, + &token.Token{ID: token.T_WHITESPACE, Value: []byte("\n" + indent)}, + ) +} + +// hasDocParam reports whether a doc comment already has a @param line for the +// named parameter, whatever its type or description. +func hasDocParam(doc, name string) bool { + return regexp.MustCompile(`@param\b[^\n]*\$` + regexp.QuoteMeta(name) + `\b`).MatchString(doc) +} + +// appendDocTags inserts tag lines right before the closing `*/`, turning a +// single-line doc comment into a multi-line one first. +func appendDocTags(doc string, tags []string, indent string) string { + if !strings.Contains(doc, "\n") { + content := strings.TrimSpace(strings.TrimSuffix(strings.TrimPrefix(doc, "/**"), "*/")) + doc = "/**\n" + if content != "" { + doc += indent + " * " + content + "\n" + } + doc += indent + " */" + } + + closing := strings.LastIndex(doc, "\n") + var lines strings.Builder + for _, tag := range tags { + lines.WriteString("\n" + indent + " * " + tag) + } + return doc[:closing] + lines.String() + doc[closing:] +} + +// indentation returns the whitespace after the last line break in the leading +// tokens of a function or method, which is its indentation. +func indentation(tokens []*token.Token) string { + for i := len(tokens) - 1; i >= 0; i-- { + if tokens[i].ID != token.T_WHITESPACE { + continue + } + value := string(tokens[i].Value) + if index := strings.LastIndex(value, "\n"); index >= 0 { + return value[index+1:] + } + } + return "" +} + // leadingToken returns the first token of a function or method, which carries // the doc comment in its leading free-floating tokens. func leadingToken(node ast.Vertex) *token.Token { From 34730d3752107e40898e348eee7215e78a7d4fc2 Mon Sep 17 00:00:00 2001 From: Tomas Votruba Date: Fri, 25 Sep 2026 12:00:57 +0200 Subject: [PATCH 2/2] Use slices.Backward for indentation lookup --- internal/phpast/phpast.go | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/internal/phpast/phpast.go b/internal/phpast/phpast.go index 54c48e38..a3a21371 100644 --- a/internal/phpast/phpast.go +++ b/internal/phpast/phpast.go @@ -341,11 +341,11 @@ func appendDocTags(doc string, tags []string, indent string) string { // indentation returns the whitespace after the last line break in the leading // tokens of a function or method, which is its indentation. func indentation(tokens []*token.Token) string { - for i := len(tokens) - 1; i >= 0; i-- { - if tokens[i].ID != token.T_WHITESPACE { + for _, free := range slices.Backward(tokens) { + if free.ID != token.T_WHITESPACE { continue } - value := string(tokens[i].Value) + value := string(free.Value) if index := strings.LastIndex(value, "\n"); index >= 0 { return value[index+1:] }