Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion CLAUDE.md
Original file line number Diff line number Diff line change
Expand Up @@ -43,10 +43,13 @@ 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] [--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`:

Expand Down
20 changes: 18 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -194,6 +195,21 @@ 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
```

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
```

<br>

## How it works
Expand Down
29 changes: 17 additions & 12 deletions internal/aggregate/aggregate.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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 {
Expand All @@ -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 {
Expand Down
19 changes: 19 additions & 0 deletions internal/aggregate/aggregate_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"}},
Expand Down Expand Up @@ -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"}},
Expand Down
52 changes: 45 additions & 7 deletions internal/apply/apply.go
Original file line number Diff line number Diff line change
Expand Up @@ -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. 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)}
applier := &applier{types: types, symbols: table, inheritance: inheritance, names: phpast.ResolveNames(root), options: options}
applier.walk(root, nil)

if len(applier.added) == 0 {
Expand All @@ -36,12 +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
options Options
}

func (a *applier) walk(node ast.Vertex, class *ast.StmtClass) {
Expand Down Expand Up @@ -74,14 +104,18 @@ func (a *applier) applyFunction(function *ast.StmtFunction) {
continue
}
if resolved, ok := a.types.Function(name, position); ok {
docType := literalDocType(param, resolved)
if !a.options.allows(resolved, 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 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)
}
}
Expand Down Expand Up @@ -114,14 +148,18 @@ 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.options.allows(resolved, 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 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)
}
}
Expand Down
52 changes: 50 additions & 2 deletions internal/apply/apply_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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, apply.Options{})
return output, len(added)
}

Expand Down Expand Up @@ -324,6 +324,14 @@ func TestApply(t *testing.T) {
},
want: "<?php\nfunction pick(string $v) {}",
},
{
name: "skips string literals doc when an unresolved value is passed",
target: "<?php\nfunction pick($v) {}",
callers: []string{
"<?php\npick('a');\npick('b');\npick($value);",
},
want: "<?php\nfunction pick(string $v) {}",
},
{
name: "skips string literals doc on an already typed parameter",
target: "<?php\nfunction pick(string $v) {}",
Expand All @@ -346,8 +354,48 @@ func TestApply(t *testing.T) {

func TestNoTypesReturnsUnchanged(t *testing.T) {
src := "<?php\nfunction greet($who) {}"
output, added, changed := apply.Source([]byte(src), aggregate.Resolve(nil), symbols.New(), inherit.New())
output, added, changed := apply.Source([]byte(src), aggregate.Resolve(nil), symbols.New(), inherit.New(), apply.Options{})
if changed || len(added) != 0 || output != src {
t.Errorf("expected unchanged, got changed=%v count=%d", changed, len(added))
}
}

func TestLiteralsOnly(t *testing.T) {
src := "<?php\nfunction pick($v, $count, $page = 1) {}\npick('a', 1);\npick('b', 2);"
table := symbols.New()
types := aggregate.Resolve(collect.FromSource([]byte(src), table))

output, added, _ := apply.Source([]byte(src), types, table, inherit.New(), apply.Options{Literals: true})

want := "<?php\n/**\n * @param 'a'|'b' $v\n */\nfunction pick(string $v, $count, $page = 1) {}\npick('a', 1);\npick('b', 2);"
if output != want || len(added) != 1 {
t.Errorf("\n got: %q (%d added)\nwant: %q", output, len(added), want)
}
}

func TestObjectsOnly(t *testing.T) {
src := "<?php\nfunction save($user, $item, $count, $mixed, $status = Status::Active) {}\nenum Status {\n case Active;\n}\nsave(new User(), new Item(), 1, new Money());\nsave(null, new Order(), 2, 5);"
table := symbols.New()
table.CollectSource([]byte(src))
types := aggregate.Resolve(collect.FromSource([]byte(src), table))

output, added, _ := apply.Source([]byte(src), types, table, inherit.New(), apply.Options{Objects: true})

want := "<?php\nfunction save(?\\User $user, \\Item|\\Order $item, $count, $mixed, \\Status $status = Status::Active) {}\nenum Status {\n case Active;\n}\nsave(new User(), new Item(), 1, new Money());\nsave(null, new Order(), 2, 5);"
if output != want || len(added) != 3 {
t.Errorf("\n got: %q (%d added)\nwant: %q", output, len(added), want)
}
}

func TestLiteralsAndObjects(t *testing.T) {
src := "<?php\nfunction save($user, $mode, $count) {}\nsave(new User(), 'a', 1);\nsave(new User(), 'b', 2);"
table := symbols.New()
types := aggregate.Resolve(collect.FromSource([]byte(src), table))

output, added, _ := apply.Source([]byte(src), types, table, inherit.New(), apply.Options{Literals: true, Objects: true})

want := "<?php\n/**\n * @param 'a'|'b' $mode\n */\nfunction save(\\User $user, string $mode, $count) {}\nsave(new User(), 'a', 1);\nsave(new User(), 'b', 2);"
if output != want || len(added) != 2 {
t.Errorf("\n got: %q (%d added)\nwant: %q", output, len(added), want)
}
}
11 changes: 5 additions & 6 deletions internal/collect/collect.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@ type Record struct {
Class string // short class name for method/constructor calls
Name string // method or function name
Position int // zero-based positional argument index
Type string // "int", "float", "string", "bool", "array", "null" or "object:Fqcn"
Type string // "int", "float", "string", "bool", "array", "null" or "object:Fqcn"; empty when unresolved
IsLiteral bool // true when the argument is a plain string literal, see phpast.StringLiteral
Literal string // the string literal value, when IsLiteral
}
Expand Down Expand Up @@ -343,15 +343,14 @@ func (c *collector) record(args []ast.Vertex, base Record, sc scope) {

// argTypes returns the type(s) an argument contributes. A coalesce expression
// (`$x ?? null`) contributes the types of both sides, so `... ?? null` makes the
// parameter nullable; every other expression contributes at most one type.
// parameter nullable; every other expression contributes one type. A value that
// cannot be resolved contributes an empty type, so the aggregate knows not every
// argument was seen.
func (c *collector) argTypes(expr ast.Vertex, sc scope) []string {
if coalesce, ok := expr.(*ast.ExprBinaryCoalesce); ok {
return append(c.argTypes(coalesce.Left, sc), c.argTypes(coalesce.Right, sc)...)
}
if typeName := c.argType(expr, sc); typeName != "" {
return []string{typeName}
}
return nil
return []string{c.argType(expr, sc)}
}

// argType returns the type of an argument value as a fully qualified object
Expand Down
16 changes: 11 additions & 5 deletions internal/collect/collect_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -54,9 +54,9 @@ func TestFromSource(t *testing.T) {
want: []collect.Record{{IsFunction: true, Name: "f", Position: 0, Type: "object:Foo"}},
},
{
name: "skips variable argument",
name: "variable argument is recorded as unresolved",
src: "<?php\nf($x);",
want: nil,
want: []collect.Record{{IsFunction: true, Name: "f", Position: 0}},
},
{
name: "skips method call on untyped variable",
Expand Down Expand Up @@ -101,7 +101,10 @@ func TestFromSource(t *testing.T) {
{
name: "does not type a reassigned parameter from its new value",
src: "<?php\nclass A {\n static function fmt($day): string { return \"\"; }\n function go($d) { $d = new \\DateTime(self::fmt($d)); }\n}",
want: nil,
want: []collect.Record{
{Class: "DateTime", Name: "__construct", Position: 0},
{Class: "A", Name: "fmt", Position: 0},
},
},
{
name: "method chain rooted at new resolves to receiver class",
Expand All @@ -122,9 +125,12 @@ func TestFromSource(t *testing.T) {
},
},
{
name: "coalesce with unresolvable left contributes null only",
name: "coalesce with unresolvable left contributes unresolved and null",
src: "<?php\nclass A {\n function go($e) { $this->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",
Expand Down
Loading
Loading